Skip to main content

ntex_tls/openssl/
connect.rs

1use std::io;
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_openssl::ssl::SslConnector as OpensslConnector;
8
9use crate::{TlsConfig, openssl::SslFilter};
10
11#[derive(Clone, Debug)]
12pub struct SslConnector<S> {
13    svc: S,
14    openssl: OpensslConnector,
15}
16
17impl<A: Address> SslConnector<Connector<A>> {
18    /// Construct new `SslConnector` factory
19    pub fn new(openssl: OpensslConnector) -> Self {
20        SslConnector {
21            openssl,
22            svc: Connector::default(),
23        }
24    }
25
26    /// Use connector to open connections.
27    pub fn connector<F, S>(self, f: impl IntoService<S, SharedCfg, Connect<A>>) -> SslConnector<S>
28    where
29        S: Service<SharedCfg, Connect<A>, Res = Io<F>, Error = Error<ConnectError>>,
30    {
31        SslConnector {
32            svc: f.into_service(),
33            openssl: self.openssl,
34        }
35    }
36}
37
38impl<S> SslConnector<S> {
39    /// Establish a TLS connection on top of an existing I/O stream.
40    pub async fn connect<F: Filter>(
41        &self,
42        io: Io<F>,
43        host: &str,
44        cfg: &SharedCfg,
45    ) -> Result<Io<Layer<SslFilter, F>>, Error<ConnectError>> {
46        let cfg = cfg.get::<TlsConfig>();
47        crate::utils::connect(&cfg, host, async {
48            let ssl = self
49                .openssl
50                .configure()
51                .and_then(|config| config.into_ssl(host))
52                .map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
53            super::handshake(io, ssl, false).await
54        })
55        .await
56    }
57}
58
59impl<F: Filter, A: Address, S> Service<SharedCfg, Connect<A>> for SslConnector<S>
60where
61    S: Service<SharedCfg, Connect<A>, Res = Io<F>, Error = Error<ConnectError>>,
62{
63    type Res = Io<Layer<SslFilter, F>>;
64    type Error = Error<ConnectError>;
65
66    async fn call(
67        &self,
68        req: Connect<A>,
69        ctx: Ctx<'_, Self, SharedCfg>,
70    ) -> Result<Self::Res, Self::Error> {
71        let host = crate::utils::server_name(req.host()).to_string();
72        let io = ctx.call(&self.svc, req).await?;
73        self.connect(io, &host, ctx.st()).await
74    }
75
76    ntex_service::forward_ready!(SharedCfg, svc);
77    ntex_service::forward_shutdown!(SharedCfg, svc);
78}
79
80#[cfg(test)]
81mod tests {
82    use super::*;
83
84    use ntex_service::Pipeline;
85    use tls_openssl::ssl::SslMethod;
86
87    #[ntex::test]
88    async fn test_openssl_connect() {
89        let server = ntex::server::test_server(async || {
90            ntex::service::fn_service(async |_| Ok::<_, ()>(()))
91        });
92
93        let ssl = OpensslConnector::builder(SslMethod::tls()).unwrap();
94        let _: SslConnector<Connector<&'static str>> = SslConnector::new(ssl.build());
95        let ssl = OpensslConnector::builder(SslMethod::tls()).unwrap();
96        let svc = SslConnector::new(ssl.build()).clone();
97        assert!(format!("{svc:?}").contains("SslConnector"));
98
99        let srv = Pipeline::new(SharedCfg::default(), svc);
100        // always ready
101        assert!(srv.ready().await.is_ok());
102        let result = srv
103            .call(Connect::new("").set_addr(Some(server.addr())))
104            .await;
105        assert!(result.is_err());
106    }
107}