use super::*;
use crate::protocol::transport::poll_fd;
use crate::test_support::{flush_all, framed, io_kind, pass, read_frame, recv_frames, write_frame};
use crate::ClientError;
use gnitz_foundation::posix_io::set_sockopt_int;
use gnitz_wire::CONNECT_TIMEOUT;
use std::io::ErrorKind;
use std::net::TcpListener;
use std::sync::mpsc;
use std::time::Duration;
struct Loopback<T> {
target: String,
thread: std::thread::JoinHandle<T>,
_ca_dir: tempfile::TempDir,
}
pub(in crate::protocol::transport) type ServerEnd = rustls::StreamOwned<rustls::ServerConnection, TcpStream>;
fn handshake(sock: TcpStream, cfg: Arc<rustls::ServerConfig>) -> ServerEnd {
let mut end = rustls::StreamOwned::new(rustls::ServerConnection::new(cfg).unwrap(), sock);
while end.conn.is_handshaking() {
end.conn.complete_io(&mut end.sock).unwrap();
}
end
}
fn hello(mut end: ServerEnd) -> ServerEnd {
assert_eq!(read_frame(&mut end), gnitz_wire::HELLO);
write_frame(&mut end, &gnitz_wire::HELLO);
end
}
fn drain(mut r: impl std::io::Read) {
let _ = std::io::copy(&mut r, &mut std::io::sink());
}
impl<T: Send + 'static> Loopback<T> {
fn serve(
rcvbuf: Option<libc::c_int>,
serve: impl FnOnce(TcpStream, Arc<rustls::ServerConfig>) -> T + Send + 'static,
) -> Self {
let cert = rcgen::generate_simple_self_signed(vec!["127.0.0.1".into()]).unwrap();
let ca_dir = tempfile::tempdir().unwrap();
let pem = ca_dir.path().join("ca.pem");
std::fs::write(&pem, cert.cert.pem()).unwrap();
let key = rustls::pki_types::PrivateKeyDer::try_from(cert.signing_key.serialize_der()).unwrap();
let mut cfg = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(vec![cert.cert.der().clone()], key)
.unwrap();
cfg.alpn_protocols = vec![ALPN_GNITZ.to_vec()];
let cfg = Arc::new(cfg);
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
if let Some(sz) = rcvbuf {
set_sockopt_int(listener.as_raw_fd(), libc::SOL_SOCKET, libc::SO_RCVBUF, sz).unwrap();
}
let port = listener.local_addr().unwrap().port();
let thread = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
sock.set_nodelay(true).unwrap();
serve(sock, cfg)
});
Loopback {
target: format!("tls://127.0.0.1:{port}?ca={}", pem.display()),
thread,
_ca_dir: ca_dir,
}
}
fn join(self) -> T {
self.thread.join().unwrap()
}
}
pub(in crate::protocol::transport) fn pair(rcvbuf: Option<libc::c_int>) -> (ClientTransport, ServerEnd) {
let lb = Loopback::serve(rcvbuf, |sock, cfg| hello(handshake(sock, cfg)));
let t = ClientTransport::connect(&lb.target, Instant::now() + CONNECT_TIMEOUT).unwrap();
(t, lb.join())
}
#[test]
fn client_config_profile() {
let cfg = build_client_config(&parse_target("127.0.0.1:1").unwrap()).expect("client config");
assert!(!cfg.enable_early_data, "client 0-RTT early data must be off");
assert!(cfg
.crypto_provider()
.cipher_suites
.iter()
.all(|s| s.version() == &rustls::version::TLS13));
}
#[test]
fn ca_file_with_an_unparsable_certificate_is_refused() {
let ca = rcgen::generate_simple_self_signed(vec!["127.0.0.1".into()]).unwrap();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("ca.pem");
let pem = format!(
"{}-----BEGIN CERTIFICATE-----\nAAAA\n-----END CERTIFICATE-----\n",
ca.cert.pem()
);
std::fs::write(&path, pem).unwrap();
let target = parse_target(&format!("127.0.0.1:1?ca={}", path.display())).unwrap();
assert!(build_client_config(&target).is_err());
}
#[test]
fn one_wakeup_drains_every_frame_rustls_holds() {
let (mut t, mut peer) = pair(None);
let mut both = framed(b"one");
both.extend(framed(b"two"));
peer.write_all(&both).unwrap();
peer.flush().unwrap();
write_frame(&mut peer, b"three");
poll_fd(t.as_raw_fd(), libc::POLLIN, None).unwrap();
let (got, end) = pass(&mut t);
end.unwrap();
assert_eq!(got, [b"one".as_slice(), b"two", b"three"]);
drop(peer);
let (got, end) = pass(&mut t);
assert!(got.is_empty());
assert_eq!(io_kind(&end), Some(ErrorKind::UnexpectedEof));
}
#[test]
fn large_reply_spanning_many_records_is_intact() {
let payload: Vec<u8> = (0u8..=255).collect::<Vec<_>>().repeat(1024);
let (mut t, mut end) = pair(None);
let p = payload.clone();
let writer = std::thread::spawn(move || write_frame(&mut end, &p));
assert_eq!(recv_frames(&mut t, 1), [payload]);
writer.join().unwrap();
}
#[test]
fn flush_can_empty_the_queue_with_ciphertext_still_pending() {
let frame: Vec<u8> = (0u8..=255).collect::<Vec<_>>().repeat(240);
let (mut t, mut end) = pair(Some(8 * 1024));
set_sockopt_int(t.as_raw_fd(), libc::SOL_SOCKET, libc::SO_SNDBUF, 8 * 1024).unwrap();
t.enqueue(frame.clone());
t.flush().unwrap();
assert_eq!(t.queued_bytes(), 0, "rustls took the whole frame");
assert!(t.wants_write(), "the queue alone is not the predicate");
let reader = std::thread::spawn(move || assert_eq!(read_frame(&mut end), frame));
flush_all(&mut t);
reader.join().unwrap();
}
#[test]
fn a_frame_larger_than_the_send_buffer_leaves_its_tail_queued() {
let big: Vec<u8> = (0u8..=255)
.collect::<Vec<_>>()
.repeat((SEND_BUFFER_BYTES + 256 * 1024) / 256);
let (mut t, mut end) = pair(None);
set_sockopt_int(t.as_raw_fd(), libc::SOL_SOCKET, libc::SO_SNDBUF, 16 * 1024).unwrap();
t.enqueue(big.clone());
t.flush().unwrap();
assert!(t.queued_bytes() > 0, "the tail is queue state, not rustls's");
let reader = std::thread::spawn(move || {
assert_eq!(read_frame(&mut end), big);
assert_eq!(read_frame(&mut end), b"after");
});
t.enqueue(b"after".to_vec());
flush_all(&mut t);
reader.join().unwrap();
}
#[test]
fn silent_peer_fails_the_connect_at_the_one_deadline() {
let silent_tcp = Loopback::serve(None, |sock, _| drain(sock));
let silent_tls = Loopback::serve(None, |sock, cfg| drain(handshake(sock, cfg)));
for lb in [silent_tcp, silent_tls] {
let d = Duration::from_millis(300);
let t0 = Instant::now();
let r = ClientTransport::connect(&lb.target, t0 + d);
assert!(matches!(
r,
Err(ClientError::Protocol(ProtocolError::IoError(ref e))) if e.kind() == ErrorKind::TimedOut
));
let took = t0.elapsed();
assert!(took >= d && took < 2 * d, "{took:?}");
lb.join();
}
}
#[test]
fn another_versions_hello_is_the_refusal_when_the_close_rides_the_same_read() {
let lb = Loopback::serve(None, |sock, cfg| {
let mut end = handshake(sock, cfg);
assert_eq!(read_frame(&mut end), gnitz_wire::HELLO);
let mut other = gnitz_wire::HELLO;
other[4] ^= 1;
end.conn.writer().write_all(&framed(&other)).unwrap();
end.conn.send_close_notify();
end.flush().unwrap();
});
match ClientTransport::connect(&lb.target, Instant::now() + CONNECT_TIMEOUT) {
Err(ClientError::Refused(f)) => assert!(f.text.contains("version mismatch"), "{}", f.text),
r => panic!("expected a version refusal, got {:?}", r.err()),
}
lb.join();
}
#[test]
fn close_notify_with_bytes_behind_it_surfaces_eof() {
let (t, mut end) = pair(None);
write_frame(&mut end, b"last");
end.conn.send_close_notify();
end.flush().unwrap();
end.sock.write_all(&[0u8; 16 * 1024]).unwrap();
drop(end);
let (done_tx, done_rx) = mpsc::channel();
let reader = std::thread::spawn(move || {
let mut t = t;
let (frames, end) = pass(&mut t);
done_tx.send((frames, io_kind(&end))).unwrap();
});
let (frames, end) = done_rx
.recv_timeout(Duration::from_secs(5))
.expect("the read past close_notify must not spin");
assert_eq!(frames, [b"last".to_vec()]);
assert_eq!(end, Some(ErrorKind::UnexpectedEof));
reader.join().unwrap();
}
#[test]
fn parse_target_accepts_all_forms() {
type Parsed<'a> = (&'a str, u16, Option<&'a str>, Option<(&'a str, &'a str)>);
let cases: [(&str, Parsed); 4] = [
("db.example.com:5433", ("db.example.com", 5433, None, None)),
(
"[::1]:65535?ca=/some/dir/cert.pem",
("::1", 65535, Some("/some/dir/cert.pem"), None),
),
("h:5?cert=/c&key=/k", ("h", 5, None, Some(("/c", "/k")))),
("h:5?ca=/x&cert=/p=q&key=/k", ("h", 5, Some("/x"), Some(("/p=q", "/k")))),
];
for (input, want) in cases {
let t = parse_target(input).unwrap();
let got = (
t.host.as_str(),
t.port,
t.ca.as_deref(),
t.client_auth.as_ref().map(|(c, k)| (c.as_str(), k.as_str())),
);
assert_eq!(got, want, "{input:?}");
}
}
#[test]
fn parse_target_rejects_malformed() {
for bad in [
"", "hostonly", ":443", "h:0x1f", "h:99999", "h:443?", "h:443?ca=", "h:443?key=", "h:443?CA=/x", "[::1]443", "[::1:443", "h:443?cert=/c", "h:443?key=/k", "h:443?cert=/c&cert=/d", "h:443?ca=/x&ca=/y", "h:443?ca=/x&", "h:443?&ca=/x", ] {
assert!(
parse_target(bad).is_err(),
"{bad:?} must be rejected by the target parser"
);
}
}