Skip to main content

ntex_tls/rustls/
client.rs

1//! TLS client filter backed by rustls
2use std::{any, cell::UnsafeCell, io, sync::Arc, task::Poll};
3
4use ntex_io::{Filter, FilterBuf, FilterLayer, Io, Layer};
5use tls_rustls::{ClientConfig, ClientConnection, pki_types::ServerName};
6
7use super::stream::Stream;
8
9#[derive(Debug)]
10/// An implementation of TLS streams
11pub struct TlsClientFilter {
12    session: UnsafeCell<ClientConnection>,
13}
14
15impl FilterLayer for TlsClientFilter {
16    fn query(&self, id: any::TypeId) -> Option<Box<dyn any::Any>> {
17        self.stream(|s| s.query(id))
18    }
19
20    fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
21        self.stream(|s| s.process_read_buf(buf))
22    }
23
24    fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
25        self.stream(|s| s.process_write_buf(buf))
26    }
27
28    fn shutdown(&self, buf: &FilterBuf<'_>) -> io::Result<Poll<()>> {
29        self.stream(|s| s.shutdown(buf))
30    }
31}
32
33impl TlsClientFilter {
34    pub async fn create<F: Filter>(
35        io: Io<F>,
36        cfg: Arc<ClientConfig>,
37        domain: ServerName<'static>,
38    ) -> Result<Io<Layer<TlsClientFilter, F>>, io::Error> {
39        let mut session = ClientConnection::new(cfg, domain).map_err(io::Error::other)?;
40        session.set_buffer_limit(Some(io.cfg().write_size().capacity()));
41        let io = io.add_filter(TlsClientFilter {
42            session: UnsafeCell::new(session),
43        });
44
45        crate::utils::handshake(&io, || Ok(io.filter().is_handshaking())).await?;
46        Ok(io)
47    }
48
49    fn is_handshaking(&self) -> bool {
50        unsafe { &*self.session.get() }.is_handshaking()
51    }
52
53    fn stream<F, R>(&self, f: F) -> R
54    where
55        F: FnOnce(&mut Stream<'_, ClientConnection>) -> R,
56    {
57        let mut s = Stream::new(unsafe { &mut *self.session.get() });
58        f(&mut s)
59    }
60}