ntex_tls/rustls/
connect.rs1use 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
11pub 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 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 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}