Skip to main content

satex_server/
server.rs

1use crate::factory::HttpServiceFactory;
2use actix_server::Server as ActixServer;
3use actix_service::ServiceFactoryExt;
4use actix_tls::accept::rustls_0_23::reexports::ServerConfig;
5use actix_tls::accept::rustls_0_23::Acceptor as TlsAcceptor;
6use http::{Request, Response};
7use hyper::body::Incoming;
8use hyper::service::Service as HyperService;
9use rustls_pki_types::{CertificateDer, PrivateKeyDer};
10use satex_core::{BoxError, Error};
11use std::fs::File;
12use std::io::BufReader;
13use std::net::ToSocketAddrs;
14
15#[derive(Debug, Clone)]
16pub enum Builder {
17    Raw(RawBuilder),
18    Tls(TlsBuilder),
19}
20
21///
22/// 普通构造器
23///
24#[derive(Debug, Clone, Default)]
25pub struct RawBuilder {
26    ///
27    /// 链接最大排队数量
28    ///
29    backlog: Option<u32>,
30
31    ///
32    /// 工作线程数量
33    ///
34    workers: Option<usize>,
35
36    ///
37    /// 每个工作线程允许的最大并发链接数量
38    ///
39    max_concurrent_connections: Option<usize>,
40}
41
42impl RawBuilder {
43    pub fn new() -> Self {
44        Self::default()
45    }
46
47    pub fn workers(mut self, workers: usize) -> Self {
48        self.workers = Some(workers);
49        self
50    }
51
52    pub fn max_concurrent_connections(mut self, max: usize) -> Self {
53        self.max_concurrent_connections = Some(max);
54        self
55    }
56
57    pub fn backlog(mut self, backlog: u32) -> Self {
58        self.backlog = Some(backlog);
59        self
60    }
61
62    pub fn tls(self) -> TlsBuilder {
63        TlsBuilder::new(self)
64    }
65
66    pub fn make_service<M>(self, make_service: M) -> Server<M> {
67        Server {
68            builder: Builder::Raw(self),
69            make_service,
70        }
71    }
72}
73
74///
75/// TLS构造器
76///
77#[derive(Debug, Clone, Default)]
78pub struct TlsBuilder {
79    raw: RawBuilder,
80    certs: Option<String>,
81    private_key: Option<String>,
82    alpn_protocols: Vec<String>,
83}
84
85impl TlsBuilder {
86    pub fn new(raw: RawBuilder) -> Self {
87        Self {
88            raw,
89            certs: None,
90            private_key: None,
91            alpn_protocols: Default::default(),
92        }
93    }
94
95    pub fn certs(mut self, cert: impl Into<String>) -> Self {
96        self.certs = Some(cert.into());
97        self
98    }
99
100    pub fn private_key(mut self, key: impl Into<String>) -> Self {
101        self.private_key = Some(key.into());
102        self
103    }
104
105    pub fn extend_alpn_protocols<I: IntoIterator<Item=P>, P: Into<String>>(
106        mut self,
107        protocols: I,
108    ) -> Self {
109        self.alpn_protocols
110            .extend(protocols.into_iter().map(Into::into));
111        self
112    }
113
114    pub fn alpn_protocols<I: IntoIterator<Item=P>, P: Into<String>>(
115        mut self,
116        protocols: I,
117    ) -> Self {
118        self.alpn_protocols = protocols.into_iter().map(Into::into).collect();
119        self
120    }
121
122    pub fn workers(mut self, workers: usize) -> Self {
123        self.raw = self.raw.workers(workers);
124        self
125    }
126
127    pub fn max_concurrent_connections(mut self, max: usize) -> Self {
128        self.raw = self.raw.max_concurrent_connections(max);
129        self
130    }
131
132    pub fn backlog(mut self, backlog: u32) -> Self {
133        self.raw = self.raw.backlog(backlog);
134        self
135    }
136
137    pub fn make_service<M>(self, make_service: M) -> Server<M> {
138        Server {
139            builder: Builder::Tls(self),
140            make_service,
141        }
142    }
143}
144
145///
146/// 服务配置
147///
148pub struct Server<F> {
149    builder: Builder,
150    make_service: F,
151}
152
153impl Server<()> {
154    pub fn builder() -> RawBuilder {
155        RawBuilder::default()
156    }
157}
158
159impl<M, S, ResBody> Server<M>
160where
161    M: HyperService<(), Response=S, Error=()> + Clone + Send + 'static,
162    S: HyperService<Request<Incoming>, Response=Response<ResBody>> + Clone + 'static,
163    S::Error: Into<BoxError>,
164    ResBody: http_body::Body + 'static,
165    ResBody::Error: Into<BoxError>,
166{
167    ///
168    /// 服务绑定地址
169    ///
170    /// # Arguments
171    ///
172    /// * `name`: 服务名称
173    /// * `server`: 服务配置
174    /// * `addrs`: 绑定的地址列表
175    ///
176    /// returns: Result<Server, Error>
177    ///
178    pub fn bind<N: AsRef<str>, A: ToSocketAddrs>(
179        self,
180        name: N,
181        addrs: A,
182    ) -> Result<ActixServer, Error> {
183        let Server {
184            builder,
185            make_service,
186        } = self;
187        let (config, tls_acceptor) = match builder {
188            Builder::Raw(builder) => (builder, None),
189            Builder::Tls(builder) => {
190                let tls_acceptor = new_tls_acceptor(&builder)?;
191                (builder.raw, Some(tls_acceptor))
192            }
193        };
194
195        let mut builder = ActixServer::build();
196        if let Some(workers) = config.workers {
197            builder = builder.workers(workers);
198        }
199        if let Some(max_concurrent_connections) = config.max_concurrent_connections {
200            builder = builder.max_concurrent_connections(max_concurrent_connections);
201        }
202        if let Some(backlog) = config.backlog {
203            builder = builder.backlog(backlog);
204        }
205        builder
206            .bind(name, addrs, move || match &tls_acceptor {
207                Some(tls_acceptor) => actix_service::boxed::factory(
208                    tls_acceptor
209                        .clone()
210                        .map_err(Error::new)
211                        .and_then(HttpServiceFactory::new(make_service.clone())),
212                ),
213                None => {
214                    actix_service::boxed::factory(HttpServiceFactory::new(make_service.clone()))
215                }
216            })
217            .map(|builder| builder.run())
218            .map_err(Error::new)
219    }
220}
221
222///
223/// 根据TLS配置创建支持HTTPS的接收器
224///
225/// # Arguments
226///
227/// * `tls`: TLS配置
228///
229/// returns: Result<Acceptor, Error>
230///
231#[inline]
232fn new_tls_acceptor(tls: &TlsBuilder) -> Result<TlsAcceptor, Error> {
233    let certs = load_certs(tls.certs.as_deref())?;
234    let private_key = load_private_key(tls.private_key.as_deref())?;
235    let mut config = ServerConfig::builder()
236        .with_no_client_auth()
237        .with_single_cert(certs, private_key)
238        .map_err(Error::new)?;
239    config.alpn_protocols = tls
240        .alpn_protocols
241        .iter()
242        .map(|item| item.as_bytes().to_vec())
243        .collect();
244    Ok(TlsAcceptor::new(config))
245}
246
247///
248///
249/// 加载TLS证书文件
250///
251/// # Arguments
252///
253/// * `tls`: TLS配置信息
254///
255/// returns: Result<Vec<CertificateDer, Global>, Error>
256///
257#[inline]
258fn load_certs(path: Option<&str>) -> Result<Vec<CertificateDer<'static>>, Error> {
259    let certs = path.ok_or_else(|| Error::new("TLS is enabled, but miss `certs`!"))?;
260    let file = File::open(certs).map_err(|e| Error::new(format!("Load TLS certs error: {}", e)))?;
261    let mut reader = BufReader::new(file);
262    rustls_pemfile::certs(&mut reader)
263        .map(|x| x.map_err(Error::new))
264        .collect()
265}
266
267///
268/// 加载TLS密钥文件
269///
270/// # Arguments
271///
272/// * `tls`: TLS配置信息
273///
274/// returns: Result<PrivateKeyDer, Error>
275///
276///
277#[inline]
278fn load_private_key(path: Option<&str>) -> Result<PrivateKeyDer<'static>, Error> {
279    let private_key = path.ok_or_else(|| Error::new("TLS is enabled, but miss `private_key`!"))?;
280    let file = File::open(private_key).map_err(Error::new)?;
281    let mut reader = BufReader::new(file);
282    rustls_pemfile::private_key(&mut reader)
283        .map_err(Error::new)
284        .and_then(|key| {
285            key.ok_or_else(|| Error::new(format!("TLS private key is invalid: {}", private_key)))
286        })
287}