use std::io::{Read, Write};
use std::os::unix::io::RawFd;
use std::os::unix::net::UnixStream;
use std::time::{Duration, Instant};
use super::*;
use gnitz_foundation::posix_io::set_sockopt_int;
fn client_with_table(target: &str) -> (GnitzClient, String, u64, Arc<Schema>) {
let mut client = GnitzClient::connect(target).expect("connect");
let (sn, tid, schema) = create_table(
&mut client,
schema_of(&[("pk", TypeCode::I64), ("a", TypeCode::I64), ("b", TypeCode::I64)]),
);
(client, sn, tid, schema)
}
fn set_small_bufs(fd: RawFd) {
for opt in [libc::SO_SNDBUF, libc::SO_RCVBUF] {
set_sockopt_int(fd, libc::SOL_SOCKET, opt, 128 * 1024).unwrap();
}
}
fn evicted_within(s: &mut Session, ms: u64) -> bool {
let deadline = Instant::now() + Duration::from_millis(ms);
while Instant::now() < deadline {
std::thread::sleep(Duration::from_millis(200));
if s.submit(Request::AllocIds(1)).is_err() || !write_out(s) {
return true;
}
}
false
}
fn with_watchdog(secs: u64, f: impl FnOnce() + Send + 'static) {
let handle = std::thread::spawn(f);
let deadline = Instant::now() + Duration::from_secs(secs);
while !handle.is_finished() {
assert!(Instant::now() < deadline, "watchdog: test body deadlocked ({secs}s)");
std::thread::sleep(Duration::from_millis(50));
}
handle.join().unwrap();
}
#[test]
fn ca_pin_connects_and_bad_verifications_fail() {
let srv = ServerHandle::start_tls(1);
let mut pinned = GnitzClient::connect(&srv.tls_target()).expect("ca-pinned connect");
block_on(pinned.alloc_id()).unwrap();
let err = GnitzClient::connect(&format!("tls://{}", srv.tls_endpoint()))
.err()
.expect("webpki roots must reject");
assert!(
err.to_string().to_lowercase().contains("certificate"),
"expected a certificate error, got: {err}"
);
let other = ServerHandle::start_tls(1);
let wrong_ca = format!("tls://{}?ca={}", srv.tls_endpoint(), other.tls_ca_path().display());
assert!(
GnitzClient::connect(&wrong_ca).is_err(),
"a foreign CA must fail verification"
);
let addr = srv.tls_endpoint();
let mut roots = rustls::RootCertStore::empty();
let certs = {
use rustls::pki_types::pem::PemObject;
rustls::pki_types::CertificateDer::pem_file_iter(srv.tls_ca_path()).unwrap()
};
roots.add_parsable_certificates(certs.filter_map(Result::ok));
let mut cfg = rustls::ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth();
cfg.alpn_protocols = vec![b"http/1.1".to_vec()];
let name = rustls::pki_types::ServerName::try_from("127.0.0.1".to_string()).unwrap();
let conn = rustls::ClientConnection::new(Arc::new(cfg), name).unwrap();
let sock = std::net::TcpStream::connect(&addr).unwrap();
sock.set_read_timeout(Some(Duration::from_secs(10))).unwrap();
let mut tls = rustls::StreamOwned::new(conn, sock);
let mut failed = false;
while tls.conn.is_handshaking() {
if tls.conn.complete_io(&mut tls.sock).is_err() {
failed = true;
break;
}
}
assert!(failed, "ALPN mismatch must fail the handshake");
}
#[test]
fn ipv6_loopback_connect_and_ca_verify() {
let srv = ServerHandle::start_tls_v6(1);
let target = srv.tls_target();
assert!(target.starts_with("tls://[::1]:"), "v6 endpoint expected, got {target}");
let mut client = GnitzClient::connect(&target).expect("ca-pinned tls over [::1]");
block_on(client.alloc_id()).unwrap();
}
#[test]
fn an_mtls_server_admits_only_a_leaf_its_client_ca_signed() {
let srv = ServerHandle::start_mtls(1);
let (mut client, _, tid, schema) = client_with_table(&srv.mtls_target());
let pushed = rows(&schema, 0..500);
block_on(client.push(tid, &schema, &pushed, WireConflictMode::Update)).unwrap();
assert_eq!(
weighted_rows(&scan_all(&mut client, tid, &schema)),
weighted_rows(&pushed)
);
assert!(
GnitzClient::connect(&srv.tls_target()).is_err(),
"a client presenting no certificate"
);
let other = ServerHandle::start_mtls(1);
let (foreign_cert, foreign_key) = other.mtls_client_cert_key();
let foreign = format!(
"{}&cert={}&key={}",
srv.tls_target(),
foreign_cert.display(),
foreign_key.display()
);
assert!(
GnitzClient::connect(&foreign).is_err(),
"a client cert from an untrusted CA"
);
}
#[test]
fn a_wire_version_mismatch_is_answered_with_the_servers_hello_then_closed() {
let srv = ServerHandle::start();
let mut sock = UnixStream::connect(srv.sock_path()).unwrap();
sock.set_read_timeout(Some(Duration::from_secs(10))).unwrap();
let mut hello = gnitz_wire::HELLO;
hello[4] ^= 1;
sock.write_all(&gnitz_wire::frame_len_prefix(hello.len())).unwrap();
sock.write_all(&hello).unwrap();
let mut reply = Vec::new();
sock.read_to_end(&mut reply).unwrap();
let want = [&gnitz_wire::frame_len_prefix(hello.len())[..], &gnitz_wire::HELLO].concat();
assert_eq!(reply, want);
}
const HELLO_DEADLINE: Duration = Duration::from_millis(500);
fn assert_reaped(transport: &str, sock: &mut impl Read) {
use std::io::ErrorKind::{TimedOut, WouldBlock};
let t0 = Instant::now();
if let Err(e) = sock.read(&mut [0u8; 1]) {
assert!(!matches!(e.kind(), WouldBlock | TimedOut), "{transport}: never closed");
}
let elapsed = t0.elapsed();
assert!(
elapsed >= HELLO_DEADLINE * 3 / 5,
"{transport}: closed after {elapsed:?}"
);
}
#[test]
fn first_frame_deadline_reaps_silent_connections() {
let deadline_ms = HELLO_DEADLINE.as_millis().to_string();
let srv = ServerHandle::start_tls_with_env(4, &[("GNITZ_HELLO_TIMEOUT_MS", &deadline_ms)]);
let timeout = Some(HELLO_DEADLINE * 6);
let mut tcp = std::net::TcpStream::connect(srv.tls_endpoint()).expect("tcp connect");
tcp.set_read_timeout(timeout).unwrap();
assert_reaped("tls", &mut tcp);
let mut unix = UnixStream::connect(srv.sock_path()).expect("unix connect");
unix.set_read_timeout(timeout).unwrap();
assert_reaped("unix", &mut unix);
let mut client = GnitzClient::connect(&srv.tls_target()).expect("a normal client must connect");
block_on(client.alloc_id()).unwrap();
}
#[test]
fn global_connection_cap_closes_excess() {
let srv = ServerHandle::start_tls_with_env(1, &[("GNITZ_MAX_CONNS", "2")]);
let target = srv.tls_target();
let _c1 = GnitzClient::connect(&target).expect("1st connection under the cap");
let c2 = GnitzClient::connect(srv.sock_path()).expect("2nd connection under the cap");
assert!(
GnitzClient::connect(&target).is_err(),
"a connection accepted past the cap must be closed"
);
assert!(GnitzClient::connect(srv.sock_path()).is_err(), "on either transport");
drop(c2);
let deadline = Instant::now() + Duration::from_secs(5);
let mut reconnected = false;
while Instant::now() < deadline {
if GnitzClient::connect(&target).is_ok() {
reconnected = true;
break;
}
std::thread::sleep(Duration::from_millis(50));
}
assert!(reconnected, "a cap slot must free after a connection closes");
}
#[test]
fn both_transports_carry_every_cell_at_its_net_weight() {
let srv = ServerHandle::start_tls(4);
let (mut tls, _, tid, schema) = client_with_table(&srv.tls_target());
let mut unix = GnitzClient::connect(srv.sock_path()).unwrap();
block_on(tls.push(tid, &schema, rows(&schema, 0..700_000), WireConflictMode::Update)).unwrap();
let mut retract = rows(&schema, [500]);
retract.weights[0] = -1;
block_on(unix.push(tid, &schema, &retract, WireConflictMode::Update)).unwrap();
let want = weighted_rows(&rows(&schema, (0..700_000).filter(|&pk| pk != 500)));
for client in [&mut tls, &mut unix] {
assert_eq!(weighted_rows(&scan_all(client, tid, &schema)), want);
}
}
#[test]
fn a_restart_fails_the_old_connection_fast_and_admits_a_new_one() {
let mut srv = ServerHandle::start_tls(1);
let targets = [srv.sock_path().to_string(), srv.tls_target()];
let mut clients: Vec<_> = targets.iter().map(|t| client_with_table(t)).collect();
srv.restart();
for ((client, _, tid, schema), target) in clients.iter_mut().zip(&targets) {
let t0 = Instant::now();
for _ in 0..2 {
let err = block_on(client.scan_spec(*tid, &ReadSpec::all_rows(ReadBound::None), schema)).unwrap_err();
assert!(matches!(err, ClientError::ConnectionLost(_)), "{target}: {err:?}");
}
assert!(t0.elapsed() < Duration::from_secs(5), "{target}: {:?}", t0.elapsed());
block_on(GnitzClient::connect(target).unwrap().alloc_id()).unwrap();
}
}
#[test]
fn pipelined_pushes_ahead_of_scan_do_not_deadlock() {
let srv = ServerHandle::start_tls(4);
let target = srv.tls_target();
with_watchdog(120, move || {
let (_setup, _, tid, schema) = client_with_table(&target);
let mut s = Session::connect(&target).unwrap();
set_small_bufs(s.as_raw_fd());
let (n_pushes, per) = (30u64, 25_000u64);
let pushes: Vec<_> = (0..n_pushes)
.map(|i| {
let batch = rows(&schema, i * per..(i + 1) * per);
s.submit(push_req(tid, &schema, &batch)).unwrap()
})
.collect();
let all = ReadSpec::all_rows(ReadBound::None);
let scan = s.submit_scan(tid.into(), &all, &schema).unwrap();
assert!(write_out(&mut s), "the server keeps reading");
let data = drive(&mut s, scan).0.expect("the scan behind the pushes");
for push in pushes {
arrived(push).expect("a pipelined push");
}
assert_eq!(
data.batch.len() as u64,
n_pushes * per,
"every pipelined row comes back"
);
});
}
#[test]
fn a_stalled_scan_client_is_evicted_by_the_send_deadline() {
let srv = ServerHandle::start_tls_with_env(4, &[("GNITZ_CLIENT_SEND_TIMEOUT_MS", "1500")]);
for target in [srv.sock_path().to_string(), srv.tls_target()] {
with_watchdog(120, move || {
let (mut setup, _, tid, schema) = client_with_table(&target);
block_on(setup.push(tid, &schema, rows(&schema, 0..200_000), WireConflictMode::Update)).unwrap();
let mut s = Session::connect(&target).unwrap();
set_small_bufs(s.as_raw_fd());
let all = ReadSpec::all_rows(ReadBound::None);
let _unread = s.submit_scan(tid.into(), &all, &schema).unwrap();
assert!(write_out(&mut s));
assert!(evicted_within(&mut s, 8_000), "{target}: not evicted");
block_on(setup.push(tid, &schema, rows(&schema, [200_000]), WireConflictMode::Update)).unwrap();
assert_eq!(scan_all(&mut setup, tid, &schema).len(), 200_001);
});
}
}
#[test]
fn inbound_cap_breach_closes_stalled_connection() {
let srv = ServerHandle::start_tls_with_env(4, &[("GNITZ_INBOUND_MEM_BYTES", "67108864")]);
let target = srv.tls_target();
with_watchdog(120, move || {
let (mut setup, _, tid, schema) = client_with_table(&target);
block_on(setup.push(tid, &schema, rows(&schema, 0..200_000), WireConflictMode::Update)).unwrap();
let mut s = Session::connect(&target).unwrap();
set_small_bufs(s.as_raw_fd());
let all = ReadSpec::all_rows(ReadBound::None);
let _unread = s.submit_scan(tid.into(), &all, &schema).unwrap();
let flood = rows(&schema, 200_000..225_000);
let mut sent = 0;
while sent < 140 && s.submit(push_req(tid, &schema, &flood)).is_ok() && write_out(&mut s) {
sent += 1;
}
assert!(
evicted_within(&mut s, 20_000),
"server must close the cap-breaching connection (sent {sent} frames)"
);
assert_eq!(
scan_all(&mut setup, tid, &schema).len(),
200_000,
"nothing of the flood was applied"
);
});
}