Skip to main content

ntex_tls/rustls/
server.rs

1//! TLS server filter backed by rustls
2use std::{any, cell::UnsafeCell, io, sync::Arc, task::Poll};
3
4use ntex_io::{Filter, FilterBuf, FilterLayer, Io, Layer};
5use ntex_util::time::Millis;
6use tls_rustls::{ServerConfig, ServerConnection};
7
8use crate::{Servername, rustls::Stream};
9
10#[derive(Debug)]
11/// An implementation of SSL streams
12pub struct TlsServerFilter {
13    session: UnsafeCell<ServerConnection>,
14}
15
16impl FilterLayer for TlsServerFilter {
17    fn query(&self, id: any::TypeId) -> Option<Box<dyn any::Any>> {
18        self.stream(|s| {
19            if id == any::TypeId::of::<Servername>() {
20                let name = s.session.server_name()?;
21                Some(Box::new(Servername(name.to_string())) as Box<dyn any::Any>)
22            } else {
23                s.query(id)
24            }
25        })
26    }
27
28    fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
29        self.stream(|s| s.process_read_buf(buf))
30    }
31
32    fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
33        self.stream(|s| s.process_write_buf(buf))
34    }
35
36    fn shutdown(&self, buf: &FilterBuf<'_>) -> io::Result<Poll<()>> {
37        self.stream(|s| s.shutdown(buf))
38    }
39}
40
41impl TlsServerFilter {
42    pub async fn create<F: Filter>(
43        io: Io<F>,
44        cfg: Arc<ServerConfig>,
45        timeout: Millis,
46    ) -> Result<Io<Layer<TlsServerFilter, F>>, io::Error> {
47        log::trace!("{}: Initiate server connection", io.tag());
48
49        crate::utils::with_timeout(timeout, async {
50            let mut session = ServerConnection::new(cfg).map_err(io::Error::other)?;
51            session.set_buffer_limit(Some(io.cfg().write_size().capacity()));
52            let io = io.add_filter(TlsServerFilter {
53                session: UnsafeCell::new(session),
54            });
55
56            crate::utils::handshake(&io, || Ok(io.filter().is_handshaking())).await?;
57            log::trace!("{}: TLS Handshake successed", io.tag());
58            Ok(io)
59        })
60        .await
61    }
62
63    fn is_handshaking(&self) -> bool {
64        unsafe { &*self.session.get() }.is_handshaking()
65    }
66
67    fn stream<F, R>(&self, f: F) -> R
68    where
69        F: FnOnce(&mut Stream<'_, ServerConnection>) -> R,
70    {
71        let mut s = Stream::new(unsafe { &mut *self.session.get() });
72        f(&mut s)
73    }
74}