Skip to main content

ntex_tls/openssl/
mod.rs

1//! An implementation of SSL streams for ntex backed by OpenSSL
2use std::{any, borrow::ToOwned, cell::UnsafeCell, cmp, io, ptr, task::Poll};
3
4use foreign_types_shared::ForeignType;
5use ntex_bytes::{BufMut, BytePage, BytePages, BytesMut};
6use ntex_io::{Filter, FilterBuf, FilterLayer, Io, Layer, types};
7use openssl_sys as ffi;
8use tls_openssl::ssl::{self, NameType, SslStream};
9use tls_openssl::x509::X509;
10
11use crate::{PeerCertChainDer, PeerCertDer, PskIdentity, Servername};
12
13mod connect;
14pub use self::connect::SslConnector;
15
16mod accept;
17pub use self::accept::SslAcceptor;
18
19mod alloc;
20pub use self::alloc::use_global_allocator;
21
22/// Connection's peer cert
23#[derive(Debug)]
24pub struct PeerCert(pub X509);
25
26/// Connection's peer cert chain
27#[derive(Debug)]
28pub struct PeerCertChain(pub Vec<X509>);
29
30/// An implementation of SSL streams
31#[derive(Debug)]
32pub struct SslFilter {
33    inner: UnsafeCell<SslStream<IoInner>>,
34}
35
36#[derive(Debug)]
37struct IoInner {
38    source: Option<BytesMut>,
39    destination: BytePages,
40}
41
42impl io::Read for IoInner {
43    fn read(&mut self, dst: &mut [u8]) -> io::Result<usize> {
44        if let Some(ref mut buf) = self.source {
45            if buf.is_empty() {
46                Err(io::Error::from(io::ErrorKind::WouldBlock))
47            } else {
48                let len = cmp::min(buf.len(), dst.len());
49                dst[..len].copy_from_slice(&buf[..len]);
50                buf.advance_to(len);
51                Ok(len)
52            }
53        } else {
54            Err(io::Error::from(io::ErrorKind::WouldBlock))
55        }
56    }
57}
58
59impl io::Write for IoInner {
60    fn write(&mut self, src: &[u8]) -> io::Result<usize> {
61        self.destination.extend_from_slice(src);
62        Ok(src.len())
63    }
64
65    fn flush(&mut self) -> io::Result<()> {
66        Ok(())
67    }
68}
69
70impl SslFilter {
71    fn new(stream: SslStream<IoInner>) -> Self {
72        Self {
73            inner: UnsafeCell::new(stream),
74        }
75    }
76
77    fn ssl(&self) -> &ssl::SslRef {
78        // SAFETY: the filter is single-threaded, and a mutable reference to
79        // the stream exists only inside `with_buffers`, which never calls back
80        // into the filter.
81        unsafe { (*self.inner.get()).ssl() }
82    }
83
84    fn with_buffers<F, R>(&self, buf: &FilterBuf<'_>, f: F) -> R
85    where
86        F: FnOnce(&mut SslStream<IoInner>, &FilterBuf<'_>) -> R,
87    {
88        self.with_buffers_inner(buf, true, f)
89    }
90
91    /// Runs `f` for input processing.
92    ///
93    /// The current write page is not borrowed: putting it back would look
94    /// like output produced by reading, which pauses reads while the write
95    /// buffer is full. Output produced here, such as alerts or key updates,
96    /// is rare and small.
97    fn with_read_buffers<F, R>(&self, buf: &FilterBuf<'_>, f: F) -> R
98    where
99        F: FnOnce(&mut SslStream<IoInner>, &FilterBuf<'_>) -> R,
100    {
101        self.with_buffers_inner(buf, false, f)
102    }
103
104    fn with_buffers_inner<F, R>(&self, buf: &FilterBuf<'_>, reuse_page: bool, f: F) -> R
105    where
106        F: FnOnce(&mut SslStream<IoInner>, &FilterBuf<'_>) -> R,
107    {
108        // SAFETY: see `ssl()`. Neither the BIO callbacks nor the buffer
109        // operations below re-enter the filter.
110        let stream = unsafe { &mut *self.inner.get() };
111
112        let st = stream.get_mut();
113        st.source = buf.with_read_src(Option::take);
114
115        // get current page from destination buffer (optimization)
116        if reuse_page {
117            buf.with_write_buffers(|_, dst| st.destination.try_get_current_from(dst));
118        }
119
120        let result = f(stream, buf);
121
122        let st = stream.get_mut();
123        // an empty source goes back to the read buffer cache
124        if let Some(src) = st.source.take() {
125            buf.with_read_src(|buf| *buf = Some(src));
126        }
127
128        // copy internal buffer to write dst buffer
129        if !st.destination.is_empty() {
130            buf.with_write_buffers(|_, dst| {
131                st.destination.move_to(dst);
132            });
133        }
134        result
135    }
136}
137
138impl FilterLayer for SslFilter {
139    fn query(&self, id: any::TypeId) -> Option<Box<dyn any::Any>> {
140        if id == any::TypeId::of::<types::HttpProtocol>() {
141            let alpn = self.ssl().selected_alpn_protocol();
142            Some(Box::new(crate::utils::http_protocol(alpn)))
143        } else if id == any::TypeId::of::<PeerCert>() {
144            Some(Box::new(PeerCert(self.ssl().peer_certificate()?)))
145        } else if id == any::TypeId::of::<PeerCertChain>() {
146            let chain = self.ssl().peer_cert_chain()?;
147            Some(Box::new(PeerCertChain(
148                chain.iter().map(ToOwned::to_owned).collect(),
149            )))
150        } else if id == any::TypeId::of::<PeerCertDer>() {
151            let cert = self.ssl().peer_certificate()?;
152            Some(Box::new(PeerCertDer(cert.to_der().ok()?)))
153        } else if id == any::TypeId::of::<PeerCertChainDer>() {
154            let ssl = self.ssl();
155            let mut chain = Vec::new();
156            for cert in ssl.peer_cert_chain().into_iter().flatten() {
157                chain.push(cert.to_der().ok()?);
158            }
159            // the chain does not include the client's certificate on the
160            // server side
161            if let Some(cert) = ssl.peer_certificate() {
162                let cert = cert.to_der().ok()?;
163                if chain.first() != Some(&cert) {
164                    chain.insert(0, cert);
165                }
166            }
167            if chain.is_empty() {
168                None
169            } else {
170                Some(Box::new(PeerCertChainDer(chain)))
171            }
172        } else if id == any::TypeId::of::<Servername>() {
173            let name = self.ssl().servername(NameType::HOST_NAME)?;
174            Some(Box::new(Servername(name.to_string())))
175        } else if id == any::TypeId::of::<PskIdentity>() {
176            let psk_id = self.ssl().psk_identity()?;
177            Some(Box::new(PskIdentity(psk_id.to_vec())))
178        } else {
179            None
180        }
181    }
182
183    fn shutdown(&self, buf: &FilterBuf<'_>) -> io::Result<Poll<()>> {
184        let ssl_result = self.with_buffers(buf, |s, _| s.shutdown());
185        let result = match ssl_result {
186            Ok(ssl::ShutdownResult::Sent) => Ok(Poll::Pending),
187            Ok(ssl::ShutdownResult::Received) => Ok(Poll::Ready(())),
188            Err(ref e) if e.code() == ssl::ErrorCode::ZERO_RETURN => Ok(Poll::Ready(())),
189            Err(ref e)
190                if matches!(
191                    e.code(),
192                    ssl::ErrorCode::WANT_READ | ssl::ErrorCode::WANT_WRITE
193                ) =>
194            {
195                Ok(Poll::Pending)
196            }
197            Err(e) => Err(e.into_io_error().unwrap_or_else(io::Error::other)),
198        };
199
200        // Our close_notify has been sent, but the peer closed the connection
201        // without sending its own; it is never going to arrive.
202        if matches!(result, Ok(Poll::Pending)) && buf.io().is_read_eof() {
203            return Ok(Poll::Ready(()));
204        }
205        result
206    }
207
208    fn process_read_buf(&self, rb: &FilterBuf<'_>) -> io::Result<()> {
209        self.with_read_buffers(rb, |stream, buf| {
210            buf.with_read_buffers(|_, dst| {
211                loop {
212                    if dst.remaining_mut() == 0 {
213                        dst.reserve_more();
214                    }
215
216                    let chunk = dst.chunk_mut();
217                    match stream.ssl_read_uninit(chunk.as_mut()) {
218                        Ok(v) => unsafe { dst.advance_mut(v) },
219                        Err(e) => {
220                            return match e.code() {
221                                ssl::ErrorCode::WANT_READ | ssl::ErrorCode::WANT_WRITE => Ok(()),
222                                ssl::ErrorCode::ZERO_RETURN => {
223                                    rb.io().close();
224                                    Ok(())
225                                }
226                                _ => {
227                                    log::trace!("{}: SSL Error: {:?}", rb.tag(), e);
228                                    Err(io::Error::other(e))
229                                }
230                            };
231                        }
232                    }
233                }
234            })
235        })
236    }
237
238    fn process_write_buf(&self, wb: &FilterBuf<'_>) -> io::Result<()> {
239        self.with_buffers(wb, |stream, buf| {
240            buf.with_write_buffers(|w_src, _| write_pages(stream, w_src))
241        })
242    }
243}
244
245/// Maximum plaintext size of a TLS record.
246const MAX_RECORD: usize = 16 * 1024;
247
248fn write_pages(stream: &mut SslStream<IoInner>, src: &mut BytePages) -> io::Result<()> {
249    while let Some(page) = src.take() {
250        let mut page = gather(page, src);
251        match stream.ssl_write(&page) {
252            Ok(v) => {
253                page.advance_to(v);
254                src.prepend(page);
255            }
256            Err(e)
257                if matches!(
258                    e.code(),
259                    ssl::ErrorCode::WANT_READ | ssl::ErrorCode::WANT_WRITE
260                ) =>
261            {
262                // nothing is consumed, e.g. a handshake is in
263                // progress, the write is retried later
264                src.prepend(page);
265                break;
266            }
267            Err(e) => return Err(io::Error::other(e)),
268        }
269    }
270    Ok(())
271}
272
273/// Joins `page` with the following whole pages that fit into one record.
274///
275/// Every `ssl_write` call produces its own record, so a small page, e.g.
276/// response headers in front of a body, would otherwise cost a record.
277fn gather(page: BytePage, src: &mut BytePages) -> BytePage {
278    let mut buf: Option<BytesMut> = None;
279    while let Some(next) = src.take() {
280        let len = buf.as_ref().map_or(page.len(), BytesMut::len);
281        if len + next.len() > MAX_RECORD {
282            src.prepend(next);
283            break;
284        }
285        buf.get_or_insert_with(|| {
286            let mut buf = BytesMut::with_capacity(MAX_RECORD);
287            buf.extend_from_slice(&page);
288            buf
289        })
290        .extend_from_slice(&next);
291    }
292    buf.map_or(page, BytePage::from)
293}
294
295fn new_stream<F>(io: &Io<F>, ssl: ssl::Ssl) -> io::Result<SslStream<IoInner>> {
296    // Let OpenSSL pull all buffered ciphertext in one BIO read, instead of
297    // reading every record header and body separately.
298    unsafe {
299        ffi::SSL_ctrl(
300            ssl.as_ptr(),
301            ffi::SSL_CTRL_SET_READ_AHEAD,
302            1,
303            ptr::null_mut(),
304        );
305    }
306
307    let inner = IoInner {
308        source: None,
309        destination: BytePages::new(io.cfg().write_size()),
310    };
311    Ok(SslStream::new(ssl, inner)?)
312}
313
314/// Create openssl connector filter factory
315pub async fn connect<F: Filter>(
316    io: Io<F>,
317    ssl: ssl::Ssl,
318) -> Result<Io<Layer<SslFilter, F>>, io::Error> {
319    handshake(io, ssl, false).await
320}
321
322/// Add ssl filter to the io stream and drive the handshake to completion
323async fn handshake<F: Filter>(
324    io: Io<F>,
325    mut ssl: ssl::Ssl,
326    accept: bool,
327) -> io::Result<Io<Layer<SslFilter, F>>> {
328    if accept {
329        ssl.set_accept_state();
330    } else {
331        ssl.set_connect_state();
332    }
333    let stream = new_stream(&io, ssl)?;
334    let io = io.add_filter(SslFilter::new(stream));
335
336    let mut eof = false;
337    loop {
338        let result = io.with_buf(|buf| io.filter().with_buffers(buf, |s, _| s.do_handshake()))?;
339        match result {
340            Ok(()) => return Ok(io),
341            Err(e) => match e.code() {
342                ssl::ErrorCode::WANT_READ => {
343                    if eof {
344                        return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "disconnected"));
345                    }
346                    // The read that reports eof may also carry the peer's last
347                    // handshake flight, so the handshake is stepped once more
348                    // before the eof is treated as a failure.
349                    eof = io.read_notify().await?.is_none();
350                }
351                ssl::ErrorCode::WANT_WRITE => {}
352                _ => return Err(io::Error::other(e)),
353            },
354        }
355    }
356}
357
358#[cfg(test)]
359pub(crate) mod tests {
360    use ntex::codec::BytesCodec;
361    use ntex_bytes::Bytes;
362    use ntex_io::{IoConfig, testing::IoTest};
363    use ntex_service::cfg::SharedCfg;
364    use ntex_util::{future::join, time::Millis};
365    use tls_openssl::pkey::{PKey, Private};
366    use tls_openssl::ssl::{SslMethod, SslVerifyMode};
367
368    use super::*;
369
370    const CERT: &[u8] = include_bytes!("../../examples/cert.pem");
371    const KEY: &[u8] = include_bytes!("../../examples/key.pem");
372
373    fn acceptor(alpn: bool) -> ssl::SslAcceptor {
374        let mut acceptor = ssl::SslAcceptor::mozilla_intermediate(SslMethod::tls()).unwrap();
375        acceptor
376            .set_private_key(&PKey::private_key_from_pem(KEY).unwrap())
377            .unwrap();
378        acceptor
379            .set_certificate(&X509::from_pem(CERT).unwrap())
380            .unwrap();
381        if alpn {
382            acceptor.set_alpn_select_callback(|_, protos| {
383                ssl::select_next_proto(b"\x02h2", protos).ok_or(ssl::AlpnError::NOACK)
384            });
385        }
386        acceptor.build()
387    }
388
389    fn connector(alpn: bool) -> ssl::SslConnector {
390        let mut connector = ssl::SslConnector::builder(SslMethod::tls()).unwrap();
391        connector.set_verify(SslVerifyMode::NONE);
392        if alpn {
393            connector.set_alpn_protos(b"\x02h2").unwrap();
394        }
395        connector.build()
396    }
397
398    /// Leaf, intermediate and root certificates with their keys.
399    pub(crate) fn cert_chain() -> [(X509, PKey<Private>); 3] {
400        use tls_openssl::x509::{X509Builder, X509NameBuilder, extension::BasicConstraints};
401        use tls_openssl::{asn1::Asn1Time, bn::BigNum, ec, hash::MessageDigest, nid::Nid};
402
403        let issue = |cn: &str, serial: u32, ca: bool, issuer: Option<&(X509, PKey<Private>)>| {
404            let group = ec::EcGroup::from_curve_name(Nid::X9_62_PRIME256V1).unwrap();
405            let key = PKey::from_ec_key(ec::EcKey::generate(&group).unwrap()).unwrap();
406            let mut name = X509NameBuilder::new().unwrap();
407            name.append_entry_by_text("CN", cn).unwrap();
408            let name = name.build();
409            let mut builder = X509Builder::new().unwrap();
410            builder.set_version(2).unwrap();
411            let serial = BigNum::from_u32(serial).unwrap().to_asn1_integer().unwrap();
412            builder.set_serial_number(&serial).unwrap();
413            builder.set_subject_name(&name).unwrap();
414            builder.set_pubkey(&key).unwrap();
415            builder
416                .set_not_before(&Asn1Time::days_from_now(0).unwrap())
417                .unwrap();
418            builder
419                .set_not_after(&Asn1Time::days_from_now(1).unwrap())
420                .unwrap();
421            if ca {
422                let ca = BasicConstraints::new().critical().ca().build().unwrap();
423                builder.append_extension(ca).unwrap();
424            }
425            let (issuer, signer) = issuer.map_or((&*name, &key), |(c, k)| (c.subject_name(), k));
426            builder.set_issuer_name(issuer).unwrap();
427            builder.sign(signer, MessageDigest::sha256()).unwrap();
428            (builder.build(), key)
429        };
430        let root = issue("ntex root", 1, true, None);
431        let intermediate = issue("ntex intermediate", 2, true, Some(&root));
432        let leaf = issue("ntex leaf", 3, false, Some(&intermediate));
433        [leaf, intermediate, root]
434    }
435
436    fn pair() -> (IoTest, IoTest) {
437        let (client, server) = IoTest::create();
438        client.remote_buffer_cap(1 << 20);
439        server.remote_buffer_cap(1 << 20);
440        (client, server)
441    }
442
443    fn tls_cfg(timeout: Millis) -> crate::TlsConfig {
444        crate::TlsConfig {
445            handshake_timeout: timeout,
446            ..crate::TlsConfig::default()
447        }
448    }
449
450    async fn handshake_pair() -> (Io<Layer<SslFilter>>, Io<Layer<SslFilter>>) {
451        let (client, server) = pair();
452        let (client, server) = join(
453            connect(
454                Io::new(client, SharedCfg::new("CLI")),
455                connector(false)
456                    .configure()
457                    .unwrap()
458                    .into_ssl("localhost")
459                    .unwrap(),
460            ),
461            handshake(
462                Io::new(server, SharedCfg::new("SRV")),
463                ssl::Ssl::new(acceptor(false).context()).unwrap(),
464                true,
465            ),
466        )
467        .await;
468        (client.unwrap(), server.unwrap())
469    }
470
471    /// The chain starts with the peer's certificate on both sides.
472    #[ntex::test]
473    async fn peer_cert_der() {
474        let [leaf, intermediate, _] = cert_chain();
475        let der = [&leaf.0, &intermediate.0].map(|c| c.to_der().unwrap());
476
477        let mut acceptor = ssl::SslAcceptor::mozilla_intermediate(SslMethod::tls()).unwrap();
478        acceptor.set_private_key(&leaf.1).unwrap();
479        acceptor.set_certificate(&leaf.0).unwrap();
480        acceptor
481            .add_extra_chain_cert(intermediate.0.clone())
482            .unwrap();
483        acceptor.set_verify_callback(SslVerifyMode::PEER, |_, _| true);
484        let acceptor = acceptor.build();
485
486        let mut connector = ssl::SslConnector::builder(SslMethod::tls()).unwrap();
487        connector.set_verify(SslVerifyMode::NONE);
488        connector.set_private_key(&leaf.1).unwrap();
489        connector.set_certificate(&leaf.0).unwrap();
490        connector.add_extra_chain_cert(intermediate.0).unwrap();
491        let ssl = connector
492            .build()
493            .configure()
494            .unwrap()
495            .into_ssl("localhost")
496            .unwrap();
497
498        let (client, server) = pair();
499        let (client, server) = join(
500            connect(Io::new(client, SharedCfg::new("CLI")), ssl),
501            handshake(
502                Io::new(server, SharedCfg::new("SRV")),
503                ssl::Ssl::new(acceptor.context()).unwrap(),
504                true,
505            ),
506        )
507        .await;
508        for io in [client.unwrap(), server.unwrap()] {
509            let cert = io.query::<PeerCertDer>();
510            assert_eq!(cert.as_ref().map(|c| &c.0), Some(&der[0]));
511            let chain = io.query::<PeerCertChainDer>();
512            assert_eq!(chain.as_ref().map(|c| &c.0[..]), Some(&der[..]));
513        }
514    }
515
516    #[ntex::test]
517    async fn acceptor_and_connector() {
518        use std::{cell::RefCell, rc::Rc};
519
520        use ntex_error::Error;
521        use ntex_net::connect::{Connect, ConnectError, Connector};
522        use ntex_service::{Pipeline, fn_service};
523
524        let (client, server) = pair();
525        let client = Rc::new(RefCell::new(Some(Io::new(client, SharedCfg::new("CLI")))));
526
527        let acceptor = SslAcceptor::from(acceptor(true)).clone();
528        assert!(format!("{acceptor:?}").contains("SslAcceptor"));
529        let acceptor = Pipeline::new((), acceptor);
530
531        let connector = SslConnector::<Connector<&str>>::new(connector(true)).connector(
532            fn_service(async move |_: Connect<&str>| {
533                Ok::<_, Error<ConnectError>>(client.borrow_mut().take().unwrap())
534            }),
535        );
536        let connector = Pipeline::new(SharedCfg::new("CLI").build(), connector);
537
538        let (server, client) = join(
539            acceptor.call(Io::new(server, SharedCfg::new("SRV"))),
540            connector.call(Connect::new("localhost:443")),
541        )
542        .await;
543        let (server, client) = (server.unwrap(), client.unwrap());
544
545        assert_eq!(
546            client.query::<types::HttpProtocol>().as_ref(),
547            Some(&types::HttpProtocol::Http2)
548        );
549        assert!(client.query::<PeerCert>().as_ref().is_some());
550        assert_eq!(
551            client.query::<PeerCertChain>().as_ref().map(|c| c.0.len()),
552            Some(1)
553        );
554        assert!(client.query::<PskIdentity>().as_ref().is_none());
555        assert_eq!(
556            server.query::<Servername>().as_ref().map(|s| s.0.as_str()),
557            Some("localhost")
558        );
559        // no client auth
560        assert!(server.query::<PeerCert>().as_ref().is_none());
561        assert!(server.query::<PeerCertChain>().as_ref().is_none());
562        assert!(server.query::<PeerCertDer>().as_ref().is_none());
563        assert!(server.query::<PeerCertChainDer>().as_ref().is_none());
564        assert!(server.query::<u32>().as_ref().is_none());
565
566        // larger than the read buffer
567        let data = Bytes::from(vec![b'a'; 256 * 1024]);
568        client.send(data.clone(), &BytesCodec).await.unwrap();
569        let mut received = 0;
570        while received < data.len() {
571            received += server.recv(&BytesCodec).await.unwrap().unwrap().len();
572        }
573
574        // close_notify is exchanged in both directions
575        let (res, ()) = join(client.shutdown(), async {
576            assert!(server.recv(&BytesCodec).await.unwrap().is_none());
577        })
578        .await;
579        res.unwrap();
580    }
581
582    #[ntex::test]
583    async fn without_alpn_and_sni() {
584        let (client, server) = pair();
585        let (client, server) = join(
586            connect(
587                Io::new(client, SharedCfg::new("CLI")),
588                ssl::Ssl::new(connector(false).context()).unwrap(),
589            ),
590            handshake(
591                Io::new(server, SharedCfg::new("SRV")),
592                ssl::Ssl::new(acceptor(false).context()).unwrap(),
593                true,
594            ),
595        )
596        .await;
597        let (client, server) = (client.unwrap(), server.unwrap());
598        assert_eq!(
599            client.query::<types::HttpProtocol>().as_ref(),
600            Some(&types::HttpProtocol::Http1)
601        );
602        assert!(server.query::<Servername>().as_ref().is_none());
603    }
604
605    #[ntex::test]
606    async fn shutdown_after_peer_disconnect() {
607        let (client, server) = handshake_pair().await;
608        // the peer goes away without close_notify
609        drop(server);
610        assert!(client.recv(&BytesCodec).await.unwrap().is_none());
611        client.shutdown().await.unwrap();
612    }
613
614    #[ntex::test]
615    async fn invalid_data_after_handshake() {
616        let (client, server) = pair();
617        let peer = server.clone();
618        let (client, server) = join(
619            connect(
620                Io::new(client, SharedCfg::new("CLI")),
621                ssl::Ssl::new(connector(false).context()).unwrap(),
622            ),
623            handshake(
624                Io::new(server, SharedCfg::new("SRV")),
625                ssl::Ssl::new(acceptor(false).context()).unwrap(),
626                true,
627            ),
628        )
629        .await;
630        let (client, _server) = (client.unwrap(), server.unwrap());
631        peer.write(b"garbage garbage garbage");
632        assert!(client.recv(&BytesCodec).await.is_err());
633    }
634
635    #[ntex::test]
636    async fn handshake_errors() {
637        use ntex_service::Pipeline;
638
639        // timeout
640        let (_client, server) = pair();
641        let io = Io::new(server, SharedCfg::new("SRV").add(tls_cfg(Millis(50))));
642        let err = Pipeline::new((), SslAcceptor::new(acceptor(false)))
643            .call(io)
644            .await
645            .unwrap_err();
646        assert_eq!(err.kind(), io::ErrorKind::TimedOut);
647
648        // peer disconnects
649        let (client, server) = pair();
650        let io = Io::new(server, SharedCfg::new("SRV"));
651        let ssl = ssl::Ssl::new(acceptor(false).context()).unwrap();
652        let (res, ()) = join(handshake(io, ssl, true), client.close()).await;
653        assert_eq!(res.unwrap_err().kind(), io::ErrorKind::UnexpectedEof);
654
655        // invalid data
656        let (client, server) = pair();
657        let io = Io::new(server, SharedCfg::new("SRV"));
658        client.write(b"GET / HTTP/1.1\r\n\r\n");
659        let ssl = ssl::Ssl::new(acceptor(false).context()).unwrap();
660        assert!(handshake(io, ssl, true).await.is_err());
661
662        // the error is reported by the connector
663        let (client, server) = pair();
664        let io = Io::new(client, SharedCfg::new("CLI"));
665        let cfg = SharedCfg::new("CLI").build();
666        let (res, ()) = join(
667            SslConnector::<ntex_net::connect::Connector<&str>>::new(connector(false)).connect(
668                io,
669                "localhost",
670                &cfg,
671            ),
672            server.close(),
673        )
674        .await;
675        assert!(res.is_err());
676    }
677
678    #[ntex::test]
679    async fn acceptor_waits_for_capacity() {
680        use ntex_service::Pipeline;
681        use ntex_util::future::lazy;
682
683        crate::MAX_SSL_ACCEPT_COUNTER.with(|c| c.set_capacity(1));
684        let acceptor = Pipeline::new((), SslAcceptor::new(acceptor(false)));
685
686        let (client, server) = pair();
687        let io = Io::new(server, SharedCfg::new("SRV").add(tls_cfg(Millis(30_000))));
688        let acceptor2 = acceptor.bind();
689        let hnd = ntex::rt::spawn(async move { acceptor2.call(io).await });
690
691        // wait until the handshake holds the only slot
692        let mut n = 0;
693        while lazy(|cx| acceptor.poll_ready(cx)).await.is_ready() {
694            n += 1;
695            assert!(n < 1000, "handshake did not start");
696            ntex_util::time::sleep(Millis(1)).await;
697        }
698        assert!(lazy(|cx| acceptor.poll_ready(cx)).await.is_pending());
699
700        // capacity is released by the failed handshake
701        client.close().await;
702        assert!(hnd.await.unwrap().is_err());
703        assert!(lazy(|cx| acceptor.poll_ready(cx)).await.is_ready());
704        crate::MAX_SSL_ACCEPT_COUNTER.with(|c| c.set_capacity(256));
705    }
706
707    /// Output written while a handshake is in progress must not be lost.
708    #[ntex::test]
709    async fn write_during_handshake_is_not_lost() {
710        let mut acceptor = ssl::SslAcceptor::mozilla_intermediate(SslMethod::tls()).unwrap();
711        acceptor
712            .set_private_key(&PKey::private_key_from_pem(KEY).unwrap())
713            .unwrap();
714        acceptor
715            .set_certificate(&X509::from_pem(CERT).unwrap())
716            .unwrap();
717        let acceptor = acceptor.build();
718        let mut connector = ssl::SslConnector::builder(SslMethod::tls()).unwrap();
719        connector.set_verify(SslVerifyMode::NONE);
720        let connector = connector.build();
721
722        let (client, server) = IoTest::create();
723        client.remote_buffer_cap(1 << 20);
724        server.remote_buffer_cap(1 << 20);
725        let server = Io::new(server, SharedCfg::new("SRV"));
726        let client = Io::new(client, SharedCfg::new("CLI"));
727
728        // the handshake is not driven, `ssl_write` starts it and reports
729        // WANT_READ
730        let mut ssl = connector
731            .configure()
732            .unwrap()
733            .into_ssl("localhost")
734            .unwrap();
735        ssl.set_connect_state();
736        let stream = new_stream(&client, ssl).unwrap();
737        let client = client.add_filter(SslFilter::new(stream));
738        client.encode_slice(b"hello").unwrap();
739
740        // reading completes the client handshake
741        ntex::rt::spawn(async move {
742            let _ = client.recv(&BytesCodec).await;
743        });
744
745        let server = handshake(server, ssl::Ssl::new(acceptor.context()).unwrap(), true)
746            .await
747            .unwrap();
748        let item = ntex_util::time::timeout(Millis(1000), server.recv(&BytesCodec))
749            .await
750            .expect("write is lost")
751            .unwrap()
752            .unwrap();
753        assert_eq!(&item[..], b"hello");
754    }
755
756    /// A drained ciphertext buffer goes back to the read buffer cache, a
757    /// dropped one would cost an allocation for every read.
758    #[ntex::test]
759    async fn drained_source_is_cached() {
760        let (client, server) = handshake_pair().await;
761
762        // the transport reads into the top of the cache, decrypted data
763        // goes to the next buffer
764        let page = server.cfg().read_size_min();
765        let get = || BytesMut::with_page_size(page);
766        let (x, y, p) = (get(), get(), get());
767        let src = p.as_ptr();
768        drop((x, y, p));
769
770        client
771            .send(Bytes::from_static(b"hello"), &BytesCodec)
772            .await
773            .unwrap();
774        let item = server.recv(&BytesCodec).await.unwrap().unwrap();
775        assert_eq!(&item[..], b"hello");
776
777        // the decrypted data buffer is released once decoded, after the source
778        let top = [get(), get()];
779        assert!(
780            top.iter().any(|b| b.as_ptr() == src),
781            "drained source is not cached"
782        );
783    }
784
785    /// Buffered output must not look like output produced by reading, that
786    /// would pause reads while the write buffer is full.
787    #[ntex::test]
788    async fn read_is_not_paused_by_buffered_output() {
789        let mut acceptor = ssl::SslAcceptor::mozilla_intermediate(SslMethod::tls()).unwrap();
790        acceptor
791            .set_private_key(&PKey::private_key_from_pem(KEY).unwrap())
792            .unwrap();
793        acceptor
794            .set_certificate(&X509::from_pem(CERT).unwrap())
795            .unwrap();
796        let acceptor = acceptor.build();
797        let mut connector = ssl::SslConnector::builder(SslMethod::tls()).unwrap();
798        connector.set_verify(SslVerifyMode::NONE);
799        let connector = connector.build();
800
801        let (client, server) = IoTest::create();
802        client.remote_buffer_cap(1 << 20);
803        server.remote_buffer_cap(1 << 20);
804        let peer = client.clone();
805
806        let server = Io::new(
807            server,
808            SharedCfg::new("SRV").add(IoConfig::default().set_write_backpressure(64)),
809        );
810        let client = Io::new(client, SharedCfg::new("CLI"));
811        let (server, client) = join(
812            handshake(server, ssl::Ssl::new(acceptor.context()).unwrap(), true),
813            connect(
814                client,
815                connector
816                    .configure()
817                    .unwrap()
818                    .into_ssl("localhost")
819                    .unwrap(),
820            ),
821        )
822        .await;
823        let (server, client) = (server.unwrap(), client.unwrap());
824
825        // the peer does not read, the output stays buffered above the high
826        // watermark
827        peer.remote_buffer_cap(0);
828        server.encode_slice(&[b'a'; 256]).unwrap();
829        // the failed write moves the output to the page list, more output
830        // leaves a partial current page
831        ntex_util::time::sleep(Millis(50)).await;
832        server.encode_slice(b"b").unwrap();
833
834        // `recv()` waits for the output to drain under write backpressure,
835        // and `read_more()` would lift a read pause, so the input is decoded
836        // as the read task delivers it
837        for msg in [&b"hello"[..], b"world", b"again"] {
838            client
839                .send(Bytes::copy_from_slice(msg), &BytesCodec)
840                .await
841                .unwrap();
842            let mut item = None;
843            for _ in 0..100 {
844                item = server.decode(&BytesCodec).unwrap();
845                if item.is_some() {
846                    break;
847                }
848                ntex_util::time::sleep(Millis(10)).await;
849            }
850            assert_eq!(&item.expect("read is paused")[..], msg);
851        }
852    }
853
854    fn stream_pair() -> (SslStream<IoInner>, SslStream<IoInner>) {
855        let stream = |ssl| {
856            let inner = IoInner {
857                source: None,
858                destination: BytePages::new(ntex_bytes::BytePageSize::Size16),
859            };
860            SslStream::new(ssl, inner).unwrap()
861        };
862        let mut ssl = connector(false)
863            .configure()
864            .unwrap()
865            .into_ssl("localhost")
866            .unwrap();
867        ssl.set_connect_state();
868        let mut client = stream(ssl);
869        let mut ssl = ssl::Ssl::new(acceptor(false).context()).unwrap();
870        ssl.set_accept_state();
871        let mut server = stream(ssl);
872
873        for _ in 0..10 {
874            let _ = client.do_handshake();
875            transfer(&mut client, &mut server);
876            let _ = server.do_handshake();
877            transfer(&mut server, &mut client);
878        }
879        assert!(client.ssl().is_init_finished() && server.ssl().is_init_finished());
880        (client, server)
881    }
882
883    fn transfer(from: &mut SslStream<IoInner>, to: &mut SslStream<IoInner>) -> Bytes {
884        let data = from.get_mut().destination.freeze();
885        let src = to.get_mut().source.get_or_insert_with(BytesMut::new);
886        src.extend_from_slice(&data);
887        data
888    }
889
890    #[allow(clippy::assert_is_empty)]
891    fn records(mut data: &[u8]) -> Vec<usize> {
892        let mut records = Vec::new();
893        while data.len() >= 5 {
894            let len = usize::from(u16::from_be_bytes([data[3], data[4]]));
895            records.push(len);
896            data = &data[5 + len..];
897        }
898        assert!(data.is_empty());
899        records
900    }
901
902    /// Small pages are joined with following pages into one record.
903    #[test]
904    fn write_gathers_pages_into_records() {
905        let (mut client, mut server) = stream_pair();
906
907        for (sizes, expected) in [
908            (&[300, 8192][..], 1),
909            (&[300, 16384][..], 2),
910            (&[100, 5000, 6000, 7000][..], 2),
911            (&[16000, 300, 84][..], 1),
912            (&[20000][..], 2),
913        ] {
914            let parts: Vec<_> = (0u8..)
915                .zip(sizes)
916                .map(|(i, &n)| Bytes::from(vec![i; n]))
917                .collect();
918            let mut src = BytePages::new(ntex_bytes::BytePageSize::Size16);
919            for p in &parts {
920                src.append(p.clone());
921            }
922            assert_eq!(src.num_pages(), parts.len());
923
924            write_pages(&mut client, &mut src).unwrap();
925            assert!(src.is_empty());
926            let wire = transfer(&mut client, &mut server);
927            assert_eq!(records(&wire).len(), expected, "{sizes:?}");
928
929            let mut plain = Vec::new();
930            let mut buf = [0u8; 4096];
931            while let Ok(n) = server.ssl_read(&mut buf) {
932                plain.extend_from_slice(&buf[..n]);
933            }
934            let expected_plain: Vec<u8> = parts.iter().flat_map(|p| p.iter().copied()).collect();
935            assert_eq!(plain, expected_plain, "{sizes:?}");
936        }
937    }
938}