Skip to main content

ntex_tls/rustls/
mod.rs

1//! An implementation of TLS streams for ntex backed by rustls
2use 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/// Connection's peer cert
18#[derive(Debug)]
19pub struct PeerCert<'a>(pub CertificateDer<'a>);
20
21/// Connection's peer cert chain
22#[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        // no client auth
216        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        // larger than the session buffer limit
223        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        // close_notify is exchanged in both directions
239        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        // no SNI for ip addresses
254        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        // the peer goes away without close_notify
261        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        // client side, the error is reported by the connector
289        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        // wait until the handshake holds the only slot
326        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        // capacity is released by the failed handshake
335        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    /// Two consecutive pages are written into one record when they fit.
342    #[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        // moves tls records between sessions, returns the records and the
351        // received plaintext
352        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}