#![recursion_limit = "256"]
#![cfg(all(windows, feature = "openssl"))]
use std::sync::{Arc, atomic::AtomicUsize, atomic::Ordering};
use ntex::client::Client;
use ntex::http::{HttpService, Version, openssl, test::server as test_server};
use ntex::service::{cfg::SharedCfg, service};
use ntex::web::{self, App, HttpResponse};
use ntex_tls::schannel::{ClientConfig, TlsConnector};
use tls_openssl::ssl::{AlpnError, SslAcceptor, SslFiletype, SslMethod};
fn ssl_acceptor() -> SslAcceptor {
ssl_acceptor_version(None)
}
fn ssl_acceptor_version(version: Option<tls_openssl::ssl::SslVersion>) -> SslAcceptor {
ssl_acceptor_builder(version).build()
}
fn ssl_acceptor_builder(
version: Option<tls_openssl::ssl::SslVersion>,
) -> tls_openssl::ssl::SslAcceptorBuilder {
let mut builder = SslAcceptor::mozilla_intermediate_v5(SslMethod::tls()).unwrap();
if version.is_some() {
builder.set_min_proto_version(version).unwrap();
builder.set_max_proto_version(version).unwrap();
}
builder
.set_private_key_file("./tests/key.pem", SslFiletype::PEM)
.unwrap();
builder
.set_certificate_chain_file("./tests/cert.pem")
.unwrap();
builder.set_alpn_select_callback(|_, protos| {
const H2: &[u8] = b"\x02h2";
if protos.windows(3).any(|window| window == H2) {
Ok(b"h2")
} else {
Err(AlpnError::NOACK)
}
});
builder.set_alpn_protos(b"\x02h2").unwrap();
builder
}
async fn join<T>(handle: std::thread::JoinHandle<T>) -> T {
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(20);
while !handle.is_finished() {
assert!(std::time::Instant::now() < deadline, "server thread hangs");
ntex::time::sleep(ntex::time::Millis(5)).await;
}
handle.join().unwrap()
}
async fn recv<T>(rx: &std::sync::mpsc::Receiver<T>) -> T {
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(20);
loop {
match rx.try_recv() {
Ok(v) => return v,
Err(std::sync::mpsc::TryRecvError::Empty) => {
assert!(std::time::Instant::now() < deadline, "server thread hangs");
ntex::time::sleep(ntex::time::Millis(5)).await;
}
Err(e) => panic!("{e}"),
}
}
}
#[ntex::test]
async fn test_connection_reuse_h2() {
let num = Arc::new(AtomicUsize::new(0));
let num2 = num.clone();
let srv = test_server(async move |_| {
let num2 = num2.clone();
service(async move |io| {
num2.fetch_add(1, Ordering::Relaxed);
Ok(io)
})
.and_then(openssl(
ssl_acceptor(),
HttpService::h2(
App::new().service(web::resource("/").route(web::to(async || HttpResponse::Ok()))),
),
))
});
let tls = TlsConnector::<ntex::connect::Connector<ntex::url::Url>>::with_config(
ClientConfig::new().danger_accept_invalid_certs(true),
);
let client = Client::builder()
.secure_connector(tls)
.build(SharedCfg::default());
let response = client.get(srv.surl("/")).send().await.unwrap();
assert!(response.status().is_success());
let response = client.post(srv.surl("/")).send().await.unwrap();
assert!(response.status().is_success());
assert_eq!(response.version(), Version::HTTP_2);
assert_eq!(num.load(Ordering::Relaxed), 1);
}
fn schannel_connector() -> TlsConnector<ntex::connect::Connector<&'static str>> {
TlsConnector::with_config(ClientConfig::new().danger_accept_invalid_certs(true))
}
#[ntex::test]
async fn test_handshake_failure_sends_alert() {
use ntex::{connect::Connect, service::Pipeline};
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(std::time::Duration::from_secs(10)))
.unwrap();
ssl_acceptor().accept(sock).unwrap_err().to_string()
});
let conn = Pipeline::new(
SharedCfg::default(),
TlsConnector::<ntex::connect::Connector<&'static str>>::with_config(ClientConfig::new()),
);
let res = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await;
assert!(res.is_err());
let err = join(server).await;
assert!(err.contains("alert unknown ca"), "no alert received: {err}");
}
#[ntex::test]
async fn test_alpn_protocols() {
use ntex::{connect::Connect, io::types::HttpProtocol, service::Pipeline};
use std::sync::Mutex;
async fn offered(config: ClientConfig) -> (Option<Vec<u8>>, HttpProtocol) {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let seen = Arc::new(Mutex::new(None));
let seen2 = seen.clone();
let (done_tx, done_rx) = std::sync::mpsc::channel::<()>();
let server = std::thread::spawn(move || {
let mut builder = SslAcceptor::mozilla_intermediate_v5(SslMethod::tls()).unwrap();
builder
.set_private_key_file("./tests/key.pem", SslFiletype::PEM)
.unwrap();
builder
.set_certificate_chain_file("./tests/cert.pem")
.unwrap();
builder.set_alpn_select_callback(move |_, protos| {
*seen2.lock().unwrap() = Some(protos.to_vec());
tls_openssl::ssl::select_next_proto(b"\x02h2\x08http/1.1", protos)
.ok_or(AlpnError::NOACK)
});
let (sock, _) = listener.accept().unwrap();
let _stream = builder.build().accept(sock).unwrap();
let _ = done_rx.recv();
});
let conn = Pipeline::new(
SharedCfg::default(),
TlsConnector::<ntex::connect::Connector<&'static str>>::with_config(
config.danger_accept_invalid_certs(true),
),
);
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
let proto = io.query::<HttpProtocol>().get().unwrap();
done_tx.send(()).unwrap();
join(server).await;
let seen = seen.lock().unwrap().take();
(seen, proto)
}
assert_eq!(
offered(ClientConfig::new()).await,
(Some(b"\x02h2\x08http/1.1".to_vec()), HttpProtocol::Http2)
);
assert_eq!(
offered(ClientConfig::new().set_alpn_protocols(&["http/1.1"])).await,
(Some(b"\x08http/1.1".to_vec()), HttpProtocol::Http1)
);
assert_eq!(
offered(ClientConfig::new().set_alpn_protocols::<&str>(&[])).await,
(None, HttpProtocol::Http1)
);
}
#[ntex::test]
async fn test_large_write_encrypted_in_one_pass() {
use ntex::{codec::BytesCodec, connect::Connect, io::Io, server, service::Pipeline};
const SIZE: usize = 16 * 1024 * 1024;
let srv = server::test_server(async || {
service(ntex::server::openssl::SslAcceptor::new(ssl_acceptor())).and_then(
async move |io: Io<_>| {
ntex::time::sleep(ntex::time::Millis(500)).await;
let mut total = 0;
while total < SIZE {
total += io.recv(&BytesCodec).await.unwrap().unwrap().len();
}
io.send(ntex::util::Bytes::from_static(b"done"), &BytesCodec)
.await
.unwrap();
Ok::<_, std::io::Error>(())
},
)
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(srv.addr())))
.await
.unwrap();
io.encode(ntex::util::Bytes::from(vec![b'x'; SIZE]), &BytesCodec)
.unwrap();
ntex::time::sleep(ntex::time::Millis(100)).await;
let pending = io
.with_buf(|buf| buf.with_write_buffers(|src, _| src.len()))
.unwrap();
assert_eq!(pending, 0, "plaintext left unencrypted after write");
io.flush(true).await.unwrap();
assert_eq!(io.recv(&BytesCodec).await.unwrap().unwrap(), "done");
}
#[ntex::test]
async fn test_records_fit_write_pages() {
use ntex::{codec::BytesCodec, connect::Connect, service::Pipeline};
use std::io::{Read, Write};
use std::sync::{Arc, Mutex};
const SIZE: usize = 256 * 1024;
#[derive(Debug)]
struct Recorder(std::net::TcpStream, Arc<Mutex<Vec<u8>>>);
impl Read for Recorder {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
let n = self.0.read(buf)?;
self.1.lock().unwrap().extend_from_slice(&buf[..n]);
Ok(n)
}
}
impl Write for Recorder {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.write(buf)
}
fn flush(&mut self) -> std::io::Result<()> {
self.0.flush()
}
}
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let raw = Arc::new(Mutex::new(Vec::new()));
let raw2 = raw.clone();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
let mut stream = ssl_acceptor().accept(Recorder(sock, raw2)).unwrap();
stream
.get_ref()
.0
.set_read_timeout(Some(std::time::Duration::from_secs(10)))
.unwrap();
let mut buf = vec![0u8; 64 * 1024];
let mut received = Vec::with_capacity(SIZE);
while received.len() < SIZE {
let n = stream.read(&mut buf).unwrap();
received.extend_from_slice(&buf[..n]);
}
stream.write_all(b"done").unwrap();
received
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
let data: Vec<u8> = (0..SIZE).map(|i| (i % 251) as u8).collect();
let (small, large) = data.split_at(SIZE / 4);
for chunk in small.chunks(997) {
io.encode(ntex::util::Bytes::copy_from_slice(chunk), &BytesCodec)
.unwrap();
if chunk[0] % 4 == 0 {
ntex::time::sleep(ntex::time::Millis(1)).await;
}
}
io.encode(ntex::util::Bytes::copy_from_slice(large), &BytesCodec)
.unwrap();
let res = ntex::time::timeout(ntex::time::Millis(15_000), io.recv(&BytesCodec)).await;
let received = join(server).await;
assert_eq!(res.unwrap().unwrap().unwrap(), "done");
assert!(received == data, "plaintext corrupted");
let raw = raw.lock().unwrap();
let mut pos = 0;
let mut records = Vec::new();
while pos + 5 <= raw.len() {
let len = u16::from_be_bytes([raw[pos + 3], raw[pos + 4]]) as usize;
if raw[pos] == 23 {
records.push(5 + len);
}
pos += 5 + len;
}
assert!(records.len() > SIZE / 16384, "{records:?}");
let page = ntex::util::BytePageSize::Size16.capacity();
assert!(
records.iter().all(|len| *len <= page),
"records larger than a write page: {records:?}"
);
}
#[ntex::test]
async fn test_shutdown_sends_close_notify() {
use ntex::{codec::BytesCodec, connect::Connect, service::Pipeline};
use std::io::{Read, Write};
use tls_openssl::ssl::ShutdownState;
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = std::sync::mpsc::channel();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
let mut stream = ssl_acceptor().accept(sock).unwrap();
stream.write_all(b"test").unwrap();
let mut buf = [0u8; 64];
let mut received = Vec::new();
let result = loop {
match stream.read(&mut buf) {
Ok(0) => break Ok(()),
Ok(n) => received.extend_from_slice(&buf[..n]),
Err(e) => break Err(e.to_string()),
}
};
let close_notify = stream.get_shutdown().contains(ShutdownState::RECEIVED);
tx.send((received, result, close_notify)).unwrap();
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
assert_eq!(io.recv(&BytesCodec).await.unwrap().unwrap(), "test");
io.encode(ntex::util::Bytes::from_static(b"bye"), &BytesCodec)
.unwrap();
io.shutdown().await.unwrap();
drop(io);
let (received, result, close_notify) = recv(&rx).await;
join(server).await;
assert_eq!(received, b"bye");
assert_eq!(result, Ok(()));
assert!(close_notify, "peer did not receive close_notify");
}
#[ntex::test]
async fn test_shutdown_waits_for_peer_close_notify() {
use ntex::{connect::Connect, service::Pipeline, time};
use std::io::Read;
use std::time::{Duration, Instant};
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = std::sync::mpsc::channel();
let (done_tx, done_rx) = std::sync::mpsc::channel::<()>();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(Duration::from_secs(10)))
.unwrap();
let mut stream = ssl_acceptor().accept(sock).unwrap();
let mut buf = [0u8; 64];
let result = stream.read(&mut buf).map_err(|e| e.to_string());
std::thread::sleep(Duration::from_millis(300));
let reply = stream.shutdown().map_err(|e| e.to_string());
tx.send((result, reply)).unwrap();
let _ = done_rx.recv_timeout(Duration::from_secs(10));
});
let cfg: SharedCfg = SharedCfg::new("CLIENT")
.add(ntex::io::IoConfig::new().set_shutdown_timeout(ntex::time::Seconds(5)))
.into();
let conn = Pipeline::new(cfg, schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
let start = Instant::now();
let res = time::timeout(Duration::from_secs(10), io.shutdown())
.await
.expect("shutdown did not complete");
let elapsed = start.elapsed();
done_tx.send(()).unwrap();
let (result, reply) = recv(&rx).await;
join(server).await;
assert_eq!(result, Ok(0), "server did not receive close_notify");
assert!(reply.is_ok(), "{reply:?}");
assert!(res.is_ok(), "{res:?}");
assert!(
elapsed >= Duration::from_millis(250),
"shutdown did not wait for the peer's close_notify: {elapsed:?}"
);
assert!(
elapsed < Duration::from_secs(3),
"the peer's close_notify was not received: {elapsed:?}"
);
drop(io);
}
async fn exchange(version: tls_openssl::ssl::SslVersion) -> String {
use ntex::{codec::BytesCodec, connect::Connect, service::Pipeline};
use std::io::{Read, Write};
use std::time::Duration;
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(Duration::from_secs(10)))
.unwrap();
let mut stream = ssl_acceptor_version(Some(version)).accept(sock).unwrap();
stream.write_all(b"hello").unwrap();
let mut buf = [0u8; 4];
stream.read_exact(&mut buf).unwrap();
assert_eq!(&buf, b"ping");
stream.write_all(b"pong").unwrap();
let mut buf = [0u8; 64];
assert_eq!(stream.read(&mut buf).unwrap(), 0, "no close_notify");
stream.shutdown().unwrap();
stream.ssl().version_str().to_string()
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
assert_eq!(io.recv(&BytesCodec).await.unwrap().unwrap(), "hello");
io.send(ntex::util::Bytes::from_static(b"ping"), &BytesCodec)
.await
.unwrap();
assert_eq!(io.recv(&BytesCodec).await.unwrap().unwrap(), "pong");
ntex::time::timeout(ntex::time::Millis(10_000), io.shutdown())
.await
.expect("shutdown timed out")
.unwrap();
join(server).await
}
#[ntex::test]
async fn test_tls13_post_handshake_messages() {
let version = exchange(tls_openssl::ssl::SslVersion::TLS1_3).await;
assert_eq!(version, "TLSv1.3");
}
#[ntex::test]
async fn test_tls12() {
let version = exchange(tls_openssl::ssl::SslVersion::TLS1_2).await;
assert_eq!(version, "TLSv1.2");
}
#[ntex::test]
async fn test_client_cert_request() {
use ntex::{connect::Connect, service::Pipeline};
use std::io::Read;
use tls_openssl::ssl::{SslVerifyMode, SslVersion};
for version in [SslVersion::TLS1_2, SslVersion::TLS1_3] {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let mut builder = ssl_acceptor_builder(Some(version));
builder.set_verify_callback(SslVerifyMode::PEER, |_, _| true);
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(std::time::Duration::from_secs(10)))
.unwrap();
let mut stream = builder.build().accept(sock).unwrap();
std::io::Write::write_all(&mut stream, b"hello").unwrap();
let mut buf = [0u8; 64];
let _ = stream.read(&mut buf);
stream.ssl().peer_certificate().is_some()
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.expect("optional client certificate request must not fail");
let data = io.recv(&ntex::codec::BytesCodec).await.unwrap().unwrap();
assert_eq!(data, "hello");
io.shutdown().await.unwrap();
assert!(!join(server).await, "client certificate was sent");
}
}
struct TestClientCert {
cert: tls_openssl::x509::X509,
pfx: Vec<u8>,
thumbprint: [u8; 20],
}
impl TestClientCert {
fn new(subject: &str, days: std::ops::Range<u32>) -> Self {
use tls_openssl::{asn1::Asn1Time, hash::MessageDigest, pkcs12::Pkcs12, pkey::PKey};
use tls_openssl::{rsa::Rsa, x509::X509Builder, x509::X509NameBuilder};
let key = PKey::from_rsa(Rsa::generate(2048).unwrap()).unwrap();
let mut name = X509NameBuilder::new().unwrap();
name.append_entry_by_text("CN", subject).unwrap();
let name = name.build();
let mut builder = X509Builder::new().unwrap();
builder.set_version(2).unwrap();
builder.set_subject_name(&name).unwrap();
builder.set_issuer_name(&name).unwrap();
builder.set_pubkey(&key).unwrap();
builder
.set_not_before(&Asn1Time::days_from_now(days.start).unwrap())
.unwrap();
builder
.set_not_after(&Asn1Time::days_from_now(days.end).unwrap())
.unwrap();
builder.sign(&key, MessageDigest::sha256()).unwrap();
let cert = builder.build();
let pfx = Pkcs12::builder()
.name(subject)
.pkey(&key)
.cert(&cert)
.build2("secret")
.unwrap()
.to_der()
.unwrap();
let thumbprint: [u8; 20] = cert
.digest(MessageDigest::sha1())
.unwrap()
.as_ref()
.try_into()
.unwrap();
Self {
cert,
pfx,
thumbprint,
}
}
fn install(&self) -> InstalledCert {
let thumbprint_hex: String = self.thumbprint.iter().map(|b| format!("{b:02X}")).collect();
let pfx_path = std::env::temp_dir().join(format!("ntex-client-{thumbprint_hex}.pfx"));
std::fs::write(&pfx_path, &self.pfx).unwrap();
let installed = InstalledCert(thumbprint_hex, pfx_path.clone());
assert!(powershell(&format!(
"Import-PfxCertificate -FilePath '{}' -CertStoreLocation Cert:\\CurrentUser\\My \
-Password (ConvertTo-SecureString secret -AsPlainText -Force) | Out-Null",
pfx_path.display()
)));
installed
}
}
#[ntex::test]
async fn test_client_cert() {
use ntex::{connect::Connect, service::Pipeline};
use std::io::Read;
use tls_openssl::ssl::{SslVerifyMode, SslVersion};
let subject = format!(
"ntex client {}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
);
let cert = TestClientCert::new(&subject, 0..2);
let earlier = loop {
let earlier = TestClientCert::new(&subject, 0..1);
if earlier.thumbprint > cert.thumbprint {
break earlier;
}
};
let not_yet_valid = TestClientCert::new(&subject, 1..3);
let _installed = [cert.install(), earlier.install(), not_yet_valid.install()];
let (cert, thumbprint) = (cert.cert, cert.thumbprint);
let client_cert = ntex_tls::schannel::Certificate::from_store(
ntex_tls::schannel::CertStoreLocation::CurrentUser,
"MY",
&thumbprint,
)
.unwrap();
assert!(client_cert.der() == cert.to_der().unwrap());
let by_subject = ntex_tls::schannel::Certificate::from_store_by_subject(
ntex_tls::schannel::CertStoreLocation::CurrentUser,
"MY",
&subject.to_uppercase(),
)
.unwrap();
assert!(by_subject.der() == cert.to_der().unwrap());
for (version, client_cert) in [
(SslVersion::TLS1_2, client_cert),
(SslVersion::TLS1_3, by_subject),
] {
let connector = TlsConnector::<ntex::connect::Connector<&'static str>>::with_config(
ClientConfig::new()
.danger_accept_invalid_certs(true)
.set_client_cert(client_cert),
);
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let mut builder = ssl_acceptor_builder(Some(version));
builder.set_verify_callback(
SslVerifyMode::PEER | SslVerifyMode::FAIL_IF_NO_PEER_CERT,
|_, _| true,
);
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(std::time::Duration::from_secs(10)))
.unwrap();
let mut stream = builder.build().accept(sock).unwrap();
std::io::Write::write_all(&mut stream, b"hello").unwrap();
let mut buf = [0u8; 64];
let _ = stream.read(&mut buf);
stream.ssl().peer_certificate().map(|c| c.to_der().unwrap())
});
let conn = Pipeline::new(SharedCfg::default(), connector);
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
let data = io.recv(&ntex::codec::BytesCodec).await.unwrap().unwrap();
assert_eq!(data, "hello");
io.shutdown().await.unwrap();
assert_eq!(
join(server).await,
Some(cert.to_der().unwrap()),
"{version:?}"
);
}
}
fn powershell(script: &str) -> bool {
std::process::Command::new("powershell")
.args(["-NoProfile", "-NonInteractive", "-Command", script])
.env_remove("PSModulePath")
.status()
.is_ok_and(|status| status.success())
}
struct InstalledCert(String, std::path::PathBuf);
impl Drop for InstalledCert {
fn drop(&mut self) {
powershell(&format!(
"Remove-Item Cert:\\CurrentUser\\My\\{} -DeleteKey",
self.0
));
let _ = std::fs::remove_file(&self.1);
}
}
#[ntex::test]
async fn test_session_resumption() {
use ntex::{connect::Connect, service::Pipeline};
use std::io::Read;
use std::time::Duration;
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = std::sync::mpsc::channel();
let server = std::thread::spawn(move || {
let acceptor = ssl_acceptor();
for _ in 0..3 {
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(Duration::from_secs(10)))
.unwrap();
let mut stream = acceptor.accept(sock).unwrap();
let reused = stream.ssl().session_reused();
std::io::Write::write_all(&mut stream, b"hello").unwrap();
let mut buf = [0u8; 64];
let _ = stream.read(&mut buf);
let _ = stream.shutdown();
tx.send(reused).unwrap();
}
});
let connector = schannel_connector();
let mut reused = Vec::new();
for _ in 0..3 {
let conn = Pipeline::new(SharedCfg::default(), connector.clone());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
let data = io.recv(&ntex::codec::BytesCodec).await.unwrap().unwrap();
assert_eq!(data, "hello");
io.shutdown().await.unwrap();
drop(io);
reused.push(recv(&rx).await);
}
join(server).await;
assert_eq!(reused, [false, true, true], "sessions are not resumed");
}
#[ntex::test]
async fn test_handshake_then_eof() {
use ntex::{connect::Connect, service::Pipeline};
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (done_tx, done_rx) = std::sync::mpsc::channel::<()>();
let server = std::thread::spawn(move || {
let acceptor = ssl_acceptor();
let mut streams = Vec::new();
while done_rx.try_recv().is_err() {
let (sock, _) = listener.accept().unwrap();
if let Ok(stream) = acceptor.accept(sock) {
stream
.get_ref()
.shutdown(std::net::Shutdown::Write)
.unwrap();
streams.push(stream);
}
}
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
for _ in 0..20 {
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.expect("handshake followed by eof must succeed");
assert!(matches!(io.recv(&ntex::codec::BytesCodec).await, Ok(None)));
}
done_tx.send(()).unwrap();
let _ = std::net::TcpStream::connect(addr);
join(server).await;
}
#[ntex::test]
async fn test_peer_close_notify_closes_io() {
use ntex::{codec::BytesCodec, connect::Connect, service::Pipeline, time};
use std::io::Read;
use std::time::Duration;
use tls_openssl::ssl::ShutdownState;
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = std::sync::mpsc::channel();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(Duration::from_secs(10)))
.unwrap();
let mut stream = ssl_acceptor().accept(sock).unwrap();
std::io::Write::write_all(&mut stream, b"test").unwrap();
stream.shutdown().unwrap();
let mut buf = [0u8; 64];
let result = stream.read(&mut buf).map_err(|e| e.to_string());
let close_notify = stream.get_shutdown().contains(ShutdownState::RECEIVED);
tx.send((result, close_notify)).unwrap();
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
assert_eq!(io.recv(&BytesCodec).await.unwrap().unwrap(), "test");
let item = time::timeout(Duration::from_secs(5), io.recv(&BytesCodec))
.await
.expect("io is not closed after peer close_notify");
assert!(matches!(item, Ok(None)), "{item:?}");
let _ = time::timeout(Duration::from_secs(5), io.shutdown()).await;
let (result, close_notify) = recv(&rx).await;
join(server).await;
assert_eq!(result, Ok(0));
assert!(close_notify, "peer did not receive close_notify");
drop(io);
}
#[derive(Debug)]
struct Trickle {
sock: std::net::TcpStream,
pending: Vec<u8>,
piece: usize,
}
impl Trickle {
fn new(sock: std::net::TcpStream) -> Self {
Self {
sock,
pending: Vec::new(),
piece: 0,
}
}
fn send(&mut self, all: bool) -> std::io::Result<()> {
const PIECES: [usize; 5] = [1, 3, 7, 100, 1000];
loop {
let piece = PIECES[self.piece % PIECES.len()];
if self.pending.is_empty() || (!all && self.pending.len() < piece) {
return Ok(());
}
let n = self.pending.len().min(piece);
std::io::Write::write_all(&mut self.sock, &self.pending[..n])?;
self.pending.drain(..n);
self.piece += 1;
std::thread::sleep(std::time::Duration::from_millis(1));
}
}
}
impl std::io::Read for Trickle {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
self.send(true)?;
self.sock.read(buf)
}
}
impl std::io::Write for Trickle {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.pending.extend_from_slice(buf);
self.send(false)?;
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
self.send(true)
}
}
#[ntex::test]
async fn test_fragmented_records() {
use ntex::{codec::BytesCodec, connect::Connect, service::Pipeline};
use std::io::{Read, Write};
use tls_openssl::ssl::SslVersion;
let payload: Vec<u8> = (0..40_000u32).map(|i| (i % 251) as u8).collect();
for version in [SslVersion::TLS1_2, SslVersion::TLS1_3] {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let data = payload.clone();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
sock.set_nodelay(true).unwrap();
sock.set_read_timeout(Some(std::time::Duration::from_secs(10)))
.unwrap();
let mut stream = ssl_acceptor_version(Some(version))
.accept(Trickle::new(sock))
.unwrap();
stream.write_all(&data).unwrap();
stream.flush().unwrap();
let mut buf = [0u8; 4];
stream.read_exact(&mut buf).unwrap();
assert_eq!(&buf, b"ping");
stream.write_all(b"pong").unwrap();
stream.flush().unwrap();
let mut buf = [0u8; 64];
assert_eq!(stream.read(&mut buf).unwrap(), 0, "no close_notify");
let _ = stream.shutdown();
let _ = stream.flush();
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
let mut received = Vec::new();
while received.len() < payload.len() {
received.extend_from_slice(&io.recv(&BytesCodec).await.unwrap().unwrap());
}
assert!(received == payload, "{version:?}");
io.send(ntex::util::Bytes::from_static(b"ping"), &BytesCodec)
.await
.unwrap();
assert_eq!(io.recv(&BytesCodec).await.unwrap().unwrap(), "pong");
io.shutdown().await.unwrap();
join(server).await;
}
}
#[derive(Debug)]
struct Coalesce {
sock: std::net::TcpStream,
pending: Vec<u8>,
}
impl Coalesce {
fn send(&mut self) -> std::io::Result<()> {
std::io::Write::write_all(&mut self.sock, &self.pending)?;
self.pending.clear();
Ok(())
}
}
impl std::io::Read for Coalesce {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
self.send()?;
self.sock.read(buf)
}
}
impl std::io::Write for Coalesce {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.pending.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
#[ntex::test]
async fn test_data_with_last_handshake_flight() {
use ntex::{codec::BytesCodec, connect::Connect, service::Pipeline};
use std::io::{Read, Write};
use tls_openssl::ssl::SslVersion;
for version in [SslVersion::TLS1_2, SslVersion::TLS1_3] {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(std::time::Duration::from_secs(10)))
.unwrap();
let sock = Coalesce {
sock,
pending: Vec::new(),
};
let mut stream = ssl_acceptor_version(Some(version)).accept(sock).unwrap();
stream.write_all(b"hello").unwrap();
stream.get_mut().send().unwrap();
let mut buf = [0u8; 64];
assert_eq!(stream.read(&mut buf).unwrap(), 0, "no close_notify");
let _ = stream.shutdown();
stream.get_mut().send().unwrap();
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
let item = ntex::time::timeout(ntex::time::Millis(5_000), io.recv(&BytesCodec))
.await
.expect("data after the handshake is not delivered");
assert_eq!(item.unwrap().unwrap(), "hello", "{version:?}");
io.shutdown().await.unwrap();
join(server).await;
}
}
#[ntex::test]
async fn test_corrupted_record_fails_read() {
use ntex::{codec::BytesCodec, connect::Connect, service::Pipeline};
use std::io::{Read, Write};
use tls_openssl::ssl::SslVersion;
for (version, mid_stream) in [
(SslVersion::TLS1_2, true),
(SslVersion::TLS1_3, true),
(SslVersion::TLS1_2, false),
(SslVersion::TLS1_3, false),
] {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (done_tx, done_rx) = std::sync::mpsc::channel::<()>();
let server = std::thread::spawn(move || {
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(std::time::Duration::from_secs(10)))
.unwrap();
let mut stream = ssl_acceptor_version(Some(version)).accept(sock).unwrap();
if mid_stream {
stream.write_all(b"hello").unwrap();
let mut buf = [0u8; 4];
stream.read_exact(&mut buf).unwrap();
assert_eq!(&buf, b"ping");
}
let mut record = vec![0x17, 0x03, 0x03, 0x00, 0x40];
record.extend((0..0x40u8).map(|i| i.wrapping_mul(37)));
stream.get_ref().write_all(&record).unwrap();
let _ = done_rx.recv_timeout(std::time::Duration::from_secs(10));
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let connect = conn.call(Connect::new("localhost").set_addr(Some(addr)));
let io = ntex::time::timeout(ntex::time::Millis(5_000), connect)
.await
.expect("connect hangs");
if mid_stream {
let io = io.unwrap();
assert_eq!(io.recv(&BytesCodec).await.unwrap().unwrap(), "hello");
io.send(ntex::util::Bytes::from_static(b"ping"), &BytesCodec)
.await
.unwrap();
let item = ntex::time::timeout(ntex::time::Millis(5_000), io.recv(&BytesCodec))
.await
.expect("corrupted record is not detected");
assert!(item.is_err(), "{version:?}: {item:?}");
} else if let Ok(io) = io {
let item = ntex::time::timeout(ntex::time::Millis(5_000), io.recv(&BytesCodec))
.await
.expect("corrupted record is not detected");
assert!(item.is_err(), "{version:?}: {item:?}");
}
done_tx.send(()).unwrap();
join(server).await;
}
}
#[ntex::test]
async fn test_tls12_renegotiation() {
use ntex::{codec::BytesCodec, connect::Connect, service::Pipeline};
use std::io::{Read, Write};
use std::time::{Duration, Instant};
use tls_openssl::ssl::{SslOptions, SslVersion};
unsafe extern "C" {
fn SSL_renegotiate(ssl: *mut std::ffi::c_void) -> std::ffi::c_int;
fn SSL_renegotiate_pending(ssl: *const std::ffi::c_void) -> std::ffi::c_int;
fn SSL_do_handshake(ssl: *mut std::ffi::c_void) -> std::ffi::c_int;
}
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let mut builder = ssl_acceptor_builder(Some(SslVersion::TLS1_2));
builder.clear_options(SslOptions::NO_RENEGOTIATION);
let (sock, _) = listener.accept().unwrap();
sock.set_read_timeout(Some(Duration::from_millis(50)))
.unwrap();
let mut stream = builder.build().accept(sock).unwrap();
stream.write_all(b"hello").unwrap();
let ssl = std::ptr::from_ref(stream.ssl()).cast_mut().cast();
assert_eq!(unsafe { SSL_renegotiate(ssl) }, 1);
assert_eq!(unsafe { SSL_do_handshake(ssl) }, 1);
let deadline = Instant::now() + Duration::from_secs(10);
let mut buf = [0u8; 64];
while unsafe { SSL_renegotiate_pending(ssl) } != 0 {
assert!(Instant::now() < deadline, "renegotiation is not completed");
match stream.read(&mut buf) {
Err(e)
if matches!(
e.kind(),
std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut
) => {}
res => panic!("unexpected read: {res:?}"),
}
}
stream.write_all(b"pong").unwrap();
stream
.get_ref()
.set_read_timeout(Some(Duration::from_secs(10)))
.unwrap();
assert_eq!(stream.read(&mut buf).unwrap(), 0, "no close_notify");
let _ = stream.shutdown();
});
let conn = Pipeline::new(SharedCfg::default(), schannel_connector());
let io = conn
.call(Connect::new("localhost").set_addr(Some(addr)))
.await
.unwrap();
assert_eq!(io.recv(&BytesCodec).await.unwrap().unwrap(), "hello");
let item = ntex::time::timeout(ntex::time::Millis(10_000), io.recv(&BytesCodec))
.await
.expect("data after renegotiation is not delivered");
assert_eq!(item.unwrap().unwrap(), "pong");
io.shutdown().await.unwrap();
join(server).await;
}