Skip to main content

ntex_tls/rustls/
connect.rs

1use std::{fmt, io, sync::Arc};
2
3use ntex_error::Error;
4use ntex_io::{Filter, Io, Layer};
5use ntex_net::connect::{Address, Connect, ConnectError, Connector};
6use ntex_service::{Ctx, IntoService, Service, cfg::SharedCfg};
7use tls_rustls::{ClientConfig, pki_types::ServerName};
8
9use crate::{TlsConfig, rustls::TlsClientFilter};
10
11/// Rustls connector factory
12pub struct TlsConnector<S> {
13    svc: S,
14    config: Arc<ClientConfig>,
15}
16
17impl<A: Address> TlsConnector<Connector<A>> {
18    pub fn new(config: ClientConfig) -> Self {
19        TlsConnector::from(Arc::new(config))
20    }
21
22    /// Use connector to open connections.
23    pub fn connector<F, S>(self, f: impl IntoService<S, SharedCfg, Connect<A>>) -> TlsConnector<S>
24    where
25        S: Service<SharedCfg, Connect<A>, Res = Io<F>, Error = Error<ConnectError>>,
26    {
27        TlsConnector {
28            svc: f.into_service(),
29            config: self.config,
30        }
31    }
32}
33
34impl<A: Address> From<Arc<ClientConfig>> for TlsConnector<Connector<A>> {
35    fn from(config: Arc<ClientConfig>) -> Self {
36        TlsConnector {
37            config,
38            svc: Connector::default(),
39        }
40    }
41}
42
43impl<'a, A: Address> From<&'a Arc<ClientConfig>> for TlsConnector<Connector<A>> {
44    fn from(config: &'a Arc<ClientConfig>) -> Self {
45        TlsConnector {
46            config: config.clone(),
47            svc: Connector::default(),
48        }
49    }
50}
51
52impl<S: Clone> Clone for TlsConnector<S> {
53    fn clone(&self) -> Self {
54        Self {
55            svc: self.svc.clone(),
56            config: self.config.clone(),
57        }
58    }
59}
60
61impl<S: fmt::Debug> fmt::Debug for TlsConnector<S> {
62    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
63        f.debug_struct("TlsConnector(rustls)")
64            .field("svc", &self.svc)
65            .finish()
66    }
67}
68
69impl<F: Filter, A: Address, S> Service<SharedCfg, Connect<A>> for TlsConnector<S>
70where
71    S: Service<SharedCfg, Connect<A>, Res = Io<F>, Error = Error<ConnectError>>,
72{
73    type Res = Io<Layer<TlsClientFilter, F>>;
74    type Error = Error<ConnectError>;
75
76    async fn call(
77        &self,
78        req: Connect<A>,
79        ctx: Ctx<'_, Self, SharedCfg>,
80    ) -> Result<Self::Res, Self::Error> {
81        let cfg = ctx.st().get::<TlsConfig>();
82        let host = crate::utils::server_name(req.host()).to_owned();
83
84        let io = ctx.call(&self.svc, req).await?;
85        crate::utils::connect(&cfg, &host, async {
86            let name = ServerName::try_from(host.as_str()).map_err(io::Error::other)?;
87            TlsClientFilter::create(io, self.config.clone(), name.to_owned()).await
88        })
89        .await
90    }
91
92    ntex_service::forward_ready!(SharedCfg, svc);
93    ntex_service::forward_shutdown!(SharedCfg, svc);
94}
95
96#[cfg(test)]
97mod tests {
98    use super::*;
99
100    use ntex_service::Pipeline;
101    use ntex_util::future::lazy;
102    use tls_rustls::RootCertStore;
103
104    #[ntex::test]
105    async fn test_rustls_connect() {
106        let server = ntex::server::test_server(async || {
107            ntex::service::fn_service(async |_| Ok::<_, ()>(()))
108        });
109
110        let cert_store = webpki_roots::TLS_SERVER_ROOTS
111            .iter()
112            .cloned()
113            .collect::<RootCertStore>();
114        let config = ClientConfig::builder()
115            .with_root_certificates(cert_store)
116            .with_no_client_auth();
117        let _: TlsConnector<Connector<&'static str>> = TlsConnector::new(config.clone()).clone();
118        let svc = TlsConnector::from(Arc::new(config)).clone();
119        assert!(format!("{svc:?}").contains("TlsConnector"), "{svc:?}");
120
121        let srv = Pipeline::new(SharedCfg::default(), svc);
122        // always ready
123        assert!(lazy(|cx| srv.poll_ready(cx)).await.is_ready());
124        let result = srv
125            .call(Connect::new("").set_addr(Some(server.addr())))
126            .await;
127        assert!(result.is_err());
128    }
129}