1use 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#[derive(Debug)]
24pub struct PeerCert(pub X509);
25
26#[derive(Debug)]
28pub struct PeerCertChain(pub Vec<X509>);
29
30#[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 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 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 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 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 if let Some(src) = st.source.take() {
125 buf.with_read_src(|buf| *buf = Some(src));
126 }
127
128 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 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 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
245const 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 src.prepend(page);
265 break;
266 }
267 Err(e) => return Err(io::Error::other(e)),
268 }
269 }
270 Ok(())
271}
272
273fn 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 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
314pub 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
322async 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 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 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 #[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 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 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 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 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 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 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 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 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 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 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 #[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 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 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 #[ntex::test]
759 async fn drained_source_is_cached() {
760 let (client, server) = handshake_pair().await;
761
762 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 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 #[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 peer.remote_buffer_cap(0);
828 server.encode_slice(&[b'a'; 256]).unwrap();
829 ntex_util::time::sleep(Millis(50)).await;
832 server.encode_slice(b"b").unwrap();
833
834 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 #[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}