//! Tests for mutual TLS configuration. use std::collections::HashMap; use std::net::SocketAddr; use std::thread; use std::time::Duration; use tidaldb::replication::WalSegmentId; use tidaldb::replication::shard::{RegionId, ShardId}; use tidaldb::replication::transport::{Transport, WalSegmentPayload}; use tidal_net::GrpcTransport; use tidal_net::config::{GrpcTransportConfig, TlsConfig}; fn free_addr() -> SocketAddr { let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); listener.local_addr().unwrap() } /// Generate a self-signed CA and server/client certificates using rcgen. fn generate_certs(dir: &std::path::Path) -> TlsConfig { use rcgen::{CertificateParams, KeyPair}; // Generate CA. let ca_key = KeyPair::generate().unwrap(); let mut ca_params = CertificateParams::new(vec!["tidaldb-ca".to_string()]).unwrap(); ca_params.is_ca = rcgen::IsCa::Ca(rcgen::BasicConstraints::Unconstrained); let ca_cert = ca_params.self_signed(&ca_key).unwrap(); // Generate server cert signed by CA. let server_key = KeyPair::generate().unwrap(); let server_params = CertificateParams::new(vec!["localhost".to_string(), "127.0.0.1".to_string()]).unwrap(); let server_cert = server_params .signed_by(&server_key, &ca_cert, &ca_key) .unwrap(); // Generate client cert signed by CA. let client_key = KeyPair::generate().unwrap(); let client_params = CertificateParams::new(vec!["tidaldb-client".to_string()]).unwrap(); let client_cert = client_params .signed_by(&client_key, &ca_cert, &ca_key) .unwrap(); // Write to files. let ca_cert_path = dir.join("ca.pem"); let server_cert_path = dir.join("server.pem"); let server_key_path = dir.join("server-key.pem"); let client_cert_path = dir.join("client.pem"); let client_key_path = dir.join("client-key.pem"); std::fs::write(&ca_cert_path, ca_cert.pem()).unwrap(); std::fs::write(&server_cert_path, server_cert.pem()).unwrap(); std::fs::write(&server_key_path, server_key.serialize_pem()).unwrap(); std::fs::write(&client_cert_path, client_cert.pem()).unwrap(); std::fs::write(&client_key_path, client_key.serialize_pem()).unwrap(); TlsConfig { ca_cert: ca_cert_path, server_cert: server_cert_path, server_key: server_key_path, client_cert: Some(client_cert_path), client_key: Some(client_key_path), } } #[test] fn mtls_send_and_receive() { let tmp = tempfile::tempdir().unwrap(); let tls = generate_certs(tmp.path()); let addr0 = free_addr(); let addr1 = free_addr(); let config0 = GrpcTransportConfig { local_shard: ShardId(0), listen_addr: addr0, peers: HashMap::from([(ShardId(1), addr1)]), tls: Some(tls.clone()), insecure: false, ..Default::default() }; let config1 = GrpcTransportConfig { local_shard: ShardId(1), listen_addr: addr1, peers: HashMap::from([(ShardId(0), addr0)]), tls: Some(tls), insecure: false, ..Default::default() }; let t0 = GrpcTransport::new(config0).expect("transport 0 with TLS"); let t1 = GrpcTransport::new(config1).expect("transport 1 with TLS"); thread::sleep(Duration::from_millis(200)); let payload = WalSegmentPayload { id: WalSegmentId::new(RegionId::SINGLE, ShardId(0), 99), bytes: vec![0xEF; 64], event_count: 2, }; t0.send_segment(ShardId(1), payload).unwrap(); let received = t1.recv_segment().unwrap(); assert_eq!(received.id.seqno, 99); assert_eq!(received.event_count, 2); } #[test] fn plaintext_rejected_when_not_insecure() { let addr = free_addr(); let config = GrpcTransportConfig { local_shard: ShardId(0), listen_addr: addr, peers: HashMap::new(), tls: None, insecure: false, // Should reject — no TLS and not insecure. ..Default::default() }; let result = GrpcTransport::new(config); assert!(result.is_err(), "should reject plaintext when not insecure"); }