1use tls_rustls::pki_types::CertificateDer;
3
4mod accept;
5mod client;
6mod connect;
7mod server;
8mod stream;
9
10pub use self::accept::TlsAcceptor;
11pub use self::client::TlsClientFilter;
12pub use self::connect::TlsConnector;
13pub use self::server::TlsServerFilter;
14
15use self::stream::Stream;
16
17#[derive(Debug)]
19pub struct PeerCert<'a>(pub CertificateDer<'a>);
20
21#[derive(Debug)]
23pub struct PeerCertChain<'a>(pub Vec<CertificateDer<'a>>);
24
25#[cfg(test)]
26mod tests {
27 use std::{cell::RefCell, io, rc::Rc, sync::Arc};
28
29 use ntex::codec::BytesCodec;
30 use ntex_bytes::Bytes;
31 use ntex_error::Error;
32 use ntex_io::{Io, Layer, testing::IoTest, types::HttpProtocol};
33 use ntex_net::connect::{Connect, ConnectError, Connector};
34 use ntex_service::{Pipeline, cfg::SharedCfg, fn_service};
35 use ntex_util::{future::join, future::lazy, time::Millis, time::sleep};
36 use tls_rustls::client::danger::{
37 HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier,
38 };
39 use tls_rustls::pki_types::{ServerName, UnixTime};
40 use tls_rustls::{ClientConfig, DigitallySignedStruct, ServerConfig, SignatureScheme};
41
42 use super::*;
43 use crate::{MAX_SSL_ACCEPT_COUNTER, Servername, TlsConfig};
44
45 const CERT: &[u8] = include_bytes!("../../examples/cert.pem");
46 const KEY: &[u8] = include_bytes!("../../examples/key.pem");
47
48 #[derive(Debug)]
49 struct NoVerify;
50
51 impl ServerCertVerifier for NoVerify {
52 fn verify_server_cert(
53 &self,
54 _: &CertificateDer<'_>,
55 _: &[CertificateDer<'_>],
56 _: &ServerName<'_>,
57 _: &[u8],
58 _: UnixTime,
59 ) -> Result<ServerCertVerified, tls_rustls::Error> {
60 Ok(ServerCertVerified::assertion())
61 }
62
63 fn verify_tls12_signature(
64 &self,
65 _: &[u8],
66 _: &CertificateDer<'_>,
67 _: &DigitallySignedStruct,
68 ) -> Result<HandshakeSignatureValid, tls_rustls::Error> {
69 Ok(HandshakeSignatureValid::assertion())
70 }
71
72 fn verify_tls13_signature(
73 &self,
74 _: &[u8],
75 _: &CertificateDer<'_>,
76 _: &DigitallySignedStruct,
77 ) -> Result<HandshakeSignatureValid, tls_rustls::Error> {
78 Ok(HandshakeSignatureValid::assertion())
79 }
80
81 fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
82 tls_rustls::crypto::ring::default_provider()
83 .signature_verification_algorithms
84 .supported_schemes()
85 }
86 }
87
88 fn server_config(alpn: bool) -> Arc<ServerConfig> {
89 let certs = rustls_pemfile::certs(&mut &CERT[..])
90 .collect::<Result<Vec<_>, _>>()
91 .unwrap();
92 let key = rustls_pemfile::private_key(&mut &KEY[..]).unwrap().unwrap();
93 let mut cfg = ServerConfig::builder_with_provider(
94 tls_rustls::crypto::ring::default_provider().into(),
95 )
96 .with_safe_default_protocol_versions()
97 .unwrap()
98 .with_no_client_auth()
99 .with_single_cert(certs, key)
100 .unwrap();
101 if alpn {
102 cfg.alpn_protocols = vec![b"h2".to_vec()];
103 }
104 Arc::new(cfg)
105 }
106
107 fn client_config(alpn: bool) -> Arc<ClientConfig> {
108 let mut cfg = ClientConfig::builder_with_provider(
109 tls_rustls::crypto::ring::default_provider().into(),
110 )
111 .with_safe_default_protocol_versions()
112 .unwrap()
113 .dangerous()
114 .with_custom_certificate_verifier(Arc::new(NoVerify))
115 .with_no_client_auth();
116 if alpn {
117 cfg.alpn_protocols = vec![b"h2".to_vec()];
118 }
119 Arc::new(cfg)
120 }
121
122 fn pair() -> (IoTest, IoTest) {
123 let (client, server) = IoTest::create();
124 client.remote_buffer_cap(1 << 20);
125 server.remote_buffer_cap(1 << 20);
126 (client, server)
127 }
128
129 fn tls_cfg(timeout: Millis) -> TlsConfig {
130 TlsConfig {
131 handshake_timeout: timeout,
132 ..TlsConfig::default()
133 }
134 }
135
136 type ClientIo = Io<Layer<TlsClientFilter>>;
137 type ServerIo = Io<Layer<TlsServerFilter>>;
138
139 async fn handshake_pair(alpn: bool, host: &str) -> (ClientIo, ServerIo) {
140 let (client, server) = pair();
141 let (client, server) = join(
142 TlsClientFilter::create(
143 Io::new(client, SharedCfg::new("CLI")),
144 client_config(alpn),
145 ServerName::try_from(host.to_string()).unwrap(),
146 ),
147 TlsServerFilter::create(
148 Io::new(server, SharedCfg::new("SRV")),
149 server_config(alpn),
150 Millis(5_000),
151 ),
152 )
153 .await;
154 (client.unwrap(), server.unwrap())
155 }
156
157 #[ntex::test]
158 async fn acceptor_and_connector() {
159 let (client, server) = pair();
160 let client = Rc::new(RefCell::new(Some(Io::new(client, SharedCfg::new("CLI")))));
161
162 let acceptor = TlsAcceptor::from(Arc::unwrap_or_clone(server_config(true)));
163 assert!(format!("{acceptor:?}").contains("TlsAcceptor"));
164 let acceptor = Pipeline::new((), acceptor);
165
166 let cfg = client_config(true);
167 let connector = TlsConnector::<Connector<&str>>::from(&cfg)
168 .connector(fn_service(async move |_: Connect<&str>| {
169 Ok::<_, Error<ConnectError>>(client.borrow_mut().take().unwrap())
170 }))
171 .clone();
172 assert!(format!("{connector:?}").contains("TlsConnector"));
173 let connector = Pipeline::new(SharedCfg::new("CLI").build(), connector);
174
175 let (server, client) = join(
176 acceptor.call(Io::new(server, SharedCfg::new("SRV"))),
177 connector.call(Connect::new("localhost:443")),
178 )
179 .await;
180 let (server, client) = (server.unwrap(), client.unwrap());
181
182 assert_eq!(
183 client.query::<HttpProtocol>().as_ref(),
184 Some(&HttpProtocol::Http2)
185 );
186 assert_eq!(
187 server.query::<HttpProtocol>().as_ref(),
188 Some(&HttpProtocol::Http2)
189 );
190 assert!(client.query::<PeerCert<'_>>().as_ref().is_some());
191 assert_eq!(
192 client
193 .query::<PeerCertChain<'_>>()
194 .as_ref()
195 .map(|c| c.0.len()),
196 Some(1)
197 );
198 let cert = client.query::<PeerCert<'_>>().as_ref().unwrap().0.to_vec();
199 assert_eq!(
200 client.query::<crate::PeerCertDer>().as_ref().map(|c| &c.0),
201 Some(&cert)
202 );
203 assert_eq!(
204 client
205 .query::<crate::PeerCertChainDer>()
206 .as_ref()
207 .map(|c| &c.0),
208 Some(&vec![cert])
209 );
210 assert!(client.query::<Servername>().as_ref().is_none());
211 assert_eq!(
212 server.query::<Servername>().as_ref().map(|s| s.0.as_str()),
213 Some("localhost")
214 );
215 assert!(server.query::<PeerCert<'_>>().as_ref().is_none());
217 assert!(server.query::<PeerCertChain<'_>>().as_ref().is_none());
218 assert!(server.query::<crate::PeerCertDer>().as_ref().is_none());
219 assert!(server.query::<crate::PeerCertChainDer>().as_ref().is_none());
220 assert!(server.query::<u32>().as_ref().is_none());
221
222 let data = Bytes::from(vec![b'a'; 256 * 1024]);
224 client.send(data.clone(), &BytesCodec).await.unwrap();
225 let mut received = 0;
226 while received < data.len() {
227 received += server.recv(&BytesCodec).await.unwrap().unwrap().len();
228 }
229 server
230 .send(Bytes::from_static(b"reply"), &BytesCodec)
231 .await
232 .unwrap();
233 assert_eq!(
234 client.recv(&BytesCodec).await.unwrap().unwrap(),
235 Bytes::from_static(b"reply")
236 );
237
238 let (res, ()) = join(client.shutdown(), async {
240 assert!(server.recv(&BytesCodec).await.unwrap().is_none());
241 })
242 .await;
243 res.unwrap();
244 }
245
246 #[ntex::test]
247 async fn without_alpn_and_sni() {
248 let (client, server) = handshake_pair(false, "127.0.0.1").await;
249 assert_eq!(
250 client.query::<HttpProtocol>().as_ref(),
251 Some(&HttpProtocol::Http1)
252 );
253 assert!(server.query::<Servername>().as_ref().is_none());
255 }
256
257 #[ntex::test]
258 async fn shutdown_after_peer_disconnect() {
259 let (client, server) = handshake_pair(false, "localhost").await;
260 drop(server);
262 assert!(client.recv(&BytesCodec).await.unwrap().is_none());
263 client.shutdown().await.unwrap();
264 }
265
266 #[ntex::test]
267 async fn handshake_timeout() {
268 let (_client, server) = pair();
269 let io = Io::new(server, SharedCfg::new("SRV").add(tls_cfg(Millis(50))));
270 let err = Pipeline::new((), TlsAcceptor::new(server_config(false)))
271 .call(io)
272 .await
273 .unwrap_err();
274 assert_eq!(err.kind(), io::ErrorKind::TimedOut);
275 }
276
277 #[ntex::test]
278 async fn handshake_disconnect() {
279 let (client, server) = pair();
280 let io = Io::new(server, SharedCfg::new("SRV"));
281 let (res, ()) = join(
282 TlsServerFilter::create(io, server_config(false), Millis::ZERO),
283 client.close(),
284 )
285 .await;
286 assert_eq!(res.unwrap_err().kind(), io::ErrorKind::UnexpectedEof);
287
288 let (client, server) = pair();
290 let client = Rc::new(RefCell::new(Some(Io::new(client, SharedCfg::new("CLI")))));
291 let connector =
292 TlsConnector::<Connector<&str>>::new(Arc::unwrap_or_clone(client_config(false)))
293 .connector(fn_service(async move |_: Connect<&str>| {
294 Ok::<_, Error<ConnectError>>(client.borrow_mut().take().unwrap())
295 }));
296 let (res, ()) = join(
297 Pipeline::new(SharedCfg::new("CLI").build(), connector).call(Connect::new("localhost")),
298 server.close(),
299 )
300 .await;
301 assert!(res.is_err());
302 }
303
304 #[ntex::test]
305 async fn handshake_invalid_data() {
306 let (client, server) = pair();
307 let io = Io::new(server, SharedCfg::new("SRV"));
308 client.write(b"GET / HTTP/1.1\r\n\r\n");
309 let err = TlsServerFilter::create(io, server_config(false), Millis::ZERO)
310 .await
311 .unwrap_err();
312 assert_eq!(err.kind(), io::ErrorKind::InvalidData);
313 }
314
315 #[ntex::test]
316 async fn acceptor_waits_for_capacity() {
317 MAX_SSL_ACCEPT_COUNTER.with(|c| c.set_capacity(1));
318 let acceptor = Pipeline::new((), TlsAcceptor::new(server_config(false)));
319
320 let (client, server) = pair();
321 let io = Io::new(server, SharedCfg::new("SRV").add(tls_cfg(Millis(30_000))));
322 let acceptor2 = acceptor.bind();
323 let hnd = ntex::rt::spawn(async move { acceptor2.call(io).await });
324
325 let mut n = 0;
327 while lazy(|cx| acceptor.poll_ready(cx)).await.is_ready() {
328 n += 1;
329 assert!(n < 1000, "handshake did not start");
330 sleep(Millis(1)).await;
331 }
332 assert!(lazy(|cx| acceptor.poll_ready(cx)).await.is_pending());
333
334 client.close().await;
336 assert!(hnd.await.unwrap().is_err());
337 assert!(lazy(|cx| acceptor.poll_ready(cx)).await.is_ready());
338 MAX_SSL_ACCEPT_COUNTER.with(|c| c.set_capacity(256));
339 }
340
341 #[test]
343 #[allow(clippy::assert_is_empty)]
344 fn write_gathers_pages_into_records() {
345 use std::io::Read;
346
347 use ntex_bytes::{BytePageSize, BytePages};
348 use tls_rustls::{ClientConnection, ServerConnection};
349
350 macro_rules! transfer {
353 ($from:expr, $to:expr) => {{
354 let mut data = Vec::new();
355 while $from.wants_write() {
356 $from.write_tls(&mut data).unwrap();
357 }
358 let mut plain = Vec::new();
359 let mut src = &data[..];
360 while !src.is_empty() {
361 $to.read_tls(&mut src).unwrap();
362 $to.process_new_packets().unwrap();
363 let _ = $to.reader().read_to_end(&mut plain);
364 }
365 (data, plain)
366 }};
367 }
368 let limit = BytePageSize::Size16.capacity();
369 let mut client = ClientConnection::new(
370 client_config(false),
371 ServerName::try_from("localhost").unwrap(),
372 )
373 .unwrap();
374 client.set_buffer_limit(Some(limit));
375 let mut server = ServerConnection::new(server_config(false)).unwrap();
376 while client.is_handshaking() || server.is_handshaking() {
377 transfer!(client, server);
378 transfer!(server, client);
379 }
380
381 for (sizes, expected) in [
382 (&[300, 8192][..], 1),
383 (&[300, limit - 300][..], 1),
384 (&[300, 16384][..], 2),
385 (&[100, 5000, 6000, 7000][..], 2),
386 (&[20000, 300][..], 2),
387 (&[20000][..], 2),
388 ] {
389 let parts: Vec<_> = (0u8..)
390 .zip(sizes)
391 .map(|(i, &n)| Bytes::from(vec![i; n]))
392 .collect();
393 let mut src = BytePages::new(BytePageSize::Size16);
394 for p in &parts {
395 src.append(p.clone());
396 }
397 assert_eq!(src.num_pages(), parts.len());
398
399 let mut dst = BytePages::new(BytePageSize::Size16);
400 stream::write_buf(&mut client, &mut src, &mut dst).unwrap();
401 assert!(src.is_empty() && !client.wants_write());
402
403 let wire = dst.freeze();
404 let mut records = 0;
405 let mut rest = &wire[..];
406 while rest.len() >= 5 {
407 rest = &rest[5 + usize::from(u16::from_be_bytes([rest[3], rest[4]]))..];
408 records += 1;
409 }
410 assert!(rest.is_empty());
411 assert_eq!(records, expected, "{sizes:?}");
412
413 let mut plain = Vec::new();
414 let mut rd = &wire[..];
415 while !rd.is_empty() {
416 server.read_tls(&mut rd).unwrap();
417 server.process_new_packets().unwrap();
418 let _ = server.reader().read_to_end(&mut plain);
419 }
420 let expected_plain: Vec<u8> = parts.iter().flat_map(|p| p.iter().copied()).collect();
421 assert_eq!(plain, expected_plain, "{sizes:?}");
422 }
423 }
424}