ntex_tls/rustls/
client.rs1use 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)]
10pub 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}