1use std::{
2 collections::HashMap,
3 future::Future,
4 net::SocketAddr,
5 pin::Pin,
6 sync::Arc,
7 task::{Context, Poll},
8};
9
10use arc_swap::ArcSwap;
11use futures::{FutureExt, Stream, future::BoxFuture, stream};
12use ipnet::IpNet;
13use openssl::{
14 ssl::{Ssl, SslAcceptor, SslMethod},
15 x509::X509,
16};
17use snafu::ResultExt;
18use tokio::{
19 io::{self, AsyncRead, AsyncWrite, ReadBuf},
20 net::{TcpListener, TcpStream},
21 sync::{OwnedSemaphorePermit, Semaphore},
22};
23use tokio_openssl::SslStream;
24use tonic::transport::{Certificate, server::Connected};
25
26use super::{
27 CreateAcceptorSnafu, HandshakeSnafu, IncomingListenerSnafu, MaybeTlsSettings, MaybeTlsStream,
28 SslBuildSnafu, TcpBindSnafu, TlsAcceptorReloader, TlsError, TlsSettings,
29};
30use crate::tcp::{self, TcpKeepaliveConfig};
31
32impl TlsSettings {
33 pub fn acceptor(&self) -> crate::tls::Result<SslAcceptor> {
34 if self.identity.is_none() {
35 Err(TlsError::MissingRequiredIdentity)
36 } else {
37 let mut acceptor =
38 SslAcceptor::mozilla_intermediate(SslMethod::tls()).context(CreateAcceptorSnafu)?;
39 self.apply_context_base(&mut acceptor, true)?;
40 Ok(acceptor.build())
41 }
42 }
43}
44
45impl MaybeTlsSettings {
46 pub async fn bind(&self, addr: &SocketAddr) -> crate::tls::Result<MaybeTlsListener> {
47 self.bind_reloadable(addr, None).await
48 }
49
50 pub async fn bind_with_allowlist(
51 &self,
52 addr: &SocketAddr,
53 allow_origin: Vec<IpNet>,
54 ) -> crate::tls::Result<MaybeTlsListener> {
55 Ok(self
56 .bind_reloadable(addr, None)
57 .await?
58 .with_allowlist(Some(allow_origin)))
59 }
60
61 pub async fn bind_reloadable(
66 &self,
67 addr: &SocketAddr,
68 reloader: Option<TlsAcceptorReloader>,
69 ) -> crate::tls::Result<MaybeTlsListener> {
70 let listener = TcpListener::bind(addr).await.context(TcpBindSnafu)?;
71
72 let acceptor = match (self, reloader) {
73 (Self::Raw(()), _) => None,
74 (Self::Tls(_), Some(reloader)) => Some(reloader.shared()),
75 (Self::Tls(tls), None) => Some(Arc::new(ArcSwap::from_pointee(tls.acceptor()?))),
76 };
77
78 Ok(MaybeTlsListener {
79 listener,
80 acceptor,
81 origin_filter: None,
82 })
83 }
84}
85
86pub struct MaybeTlsListener {
87 listener: TcpListener,
88 acceptor: Option<Arc<ArcSwap<SslAcceptor>>>,
89 origin_filter: Option<Vec<IpNet>>,
90}
91
92impl MaybeTlsListener {
93 pub async fn accept(&mut self) -> crate::tls::Result<MaybeTlsIncomingStream<TcpStream>> {
94 let listener = self
95 .listener
96 .accept()
97 .await
98 .map(|(stream, peer_addr)| {
99 let acceptor = self
100 .acceptor
101 .as_ref()
102 .map(|accptr| accptr.load().as_ref().clone());
103 MaybeTlsIncomingStream::new(stream, peer_addr, acceptor)
104 })
105 .context(IncomingListenerSnafu)?;
106
107 if let Some(origin_filter) = &self.origin_filter {
108 if origin_filter
109 .iter()
110 .any(|net| net.contains(&listener.peer_addr().ip()))
111 {
112 Ok(listener)
113 } else {
114 Err(TlsError::Connect {
115 source: std::io::ErrorKind::ConnectionRefused.into(),
116 })
117 }
118 } else {
119 Ok(listener)
120 }
121 }
122
123 async fn into_accept(
124 mut self,
125 ) -> (crate::tls::Result<MaybeTlsIncomingStream<TcpStream>>, Self) {
126 (self.accept().await, self)
127 }
128
129 pub fn accept_stream(
130 self,
131 ) -> impl Stream<Item = crate::tls::Result<MaybeTlsIncomingStream<TcpStream>>> {
132 let mut accept = Box::pin(self.into_accept());
133 stream::poll_fn(move |context| match accept.as_mut().poll(context) {
134 Poll::Ready((item, this)) => {
135 accept.set(this.into_accept());
136 Poll::Ready(Some(item))
137 }
138 Poll::Pending => Poll::Pending,
139 })
140 }
141
142 pub fn accept_stream_limited(
143 self,
144 max_connections: Option<u32>,
145 ) -> impl Stream<
146 Item = (
147 crate::tls::Result<MaybeTlsIncomingStream<TcpStream>>,
148 Option<OwnedSemaphorePermit>,
149 ),
150 > {
151 let mut connection_semaphore_future = max_connections.map(|max| {
152 let semaphore = Arc::new(Semaphore::new(max as usize));
153 let future = Box::pin(semaphore.clone().acquire_owned());
154 (semaphore, future)
155 });
156
157 let mut accept = Box::pin(self.into_accept());
158 stream::poll_fn(move |context| {
159 let permit = match connection_semaphore_future.as_mut() {
160 Some((semaphore, future)) => match future.as_mut().poll(context) {
161 Poll::Ready(permit) => {
162 future.set(semaphore.clone().acquire_owned());
163 permit.ok()
164 }
165 Poll::Pending => return Poll::Pending,
166 },
167 None => None,
168 };
169 match accept.as_mut().poll(context) {
170 Poll::Ready((item, this)) => {
171 accept.set(this.into_accept());
172 Poll::Ready(Some((item, permit)))
173 }
174 Poll::Pending => Poll::Pending,
175 }
176 })
177 }
178
179 pub fn local_addr(&self) -> Result<SocketAddr, std::io::Error> {
180 self.listener.local_addr()
181 }
182
183 #[must_use]
184 pub fn with_allowlist(mut self, allowlist: Option<Vec<IpNet>>) -> Self {
185 self.origin_filter = allowlist;
186 self
187 }
188}
189
190impl From<TcpListener> for MaybeTlsListener {
191 fn from(listener: TcpListener) -> Self {
192 Self {
193 listener,
194 acceptor: None,
195 origin_filter: None,
196 }
197 }
198}
199
200pub struct MaybeTlsIncomingStream<S> {
201 state: StreamState<S>,
202 peer_addr: SocketAddr,
206}
207
208enum StreamState<S> {
209 Accepted(MaybeTlsStream<S>),
210 Accepting(BoxFuture<'static, Result<SslStream<S>, TlsError>>),
211 AcceptError(String),
212 Closed,
213}
214
215impl<S> MaybeTlsIncomingStream<S> {
216 pub const fn peer_addr(&self) -> SocketAddr {
217 self.peer_addr
218 }
219
220 pub fn get_ref(&self) -> Option<&S> {
222 use super::MaybeTls;
223
224 match &self.state {
225 StreamState::Accepted(stream) => Some(match stream {
226 MaybeTls::Raw(s) => s,
227 MaybeTls::Tls(s) => s.get_ref(),
228 }),
229 StreamState::Accepting(_) | StreamState::AcceptError(_) | StreamState::Closed => None,
230 }
231 }
232
233 pub const fn ssl_stream(&self) -> Option<&SslStream<S>> {
234 use super::MaybeTls;
235
236 match &self.state {
237 StreamState::Accepted(stream) => match stream {
238 MaybeTls::Raw(_) => None,
239 MaybeTls::Tls(s) => Some(s),
240 },
241 StreamState::Accepting(_) | StreamState::AcceptError(_) | StreamState::Closed => None,
242 }
243 }
244
245 pub fn get_mut(&mut self) -> Option<&mut S> {
246 use super::MaybeTls;
247
248 match &mut self.state {
249 StreamState::Accepted(stream) => Some(match stream {
250 MaybeTls::Raw(s) => s,
251 MaybeTls::Tls(s) => s.get_mut(),
252 }),
253 StreamState::Accepting(_) | StreamState::AcceptError(_) | StreamState::Closed => None,
254 }
255 }
256}
257
258impl MaybeTlsIncomingStream<TcpStream> {
259 pub(super) fn new(
260 stream: TcpStream,
261 peer_addr: SocketAddr,
262 acceptor: Option<SslAcceptor>,
263 ) -> Self {
264 let state = match acceptor {
265 Some(acceptor) => StreamState::Accepting(
266 async move {
267 let ssl = Ssl::new(acceptor.context()).context(SslBuildSnafu)?;
268 let mut stream = SslStream::new(ssl, stream).context(SslBuildSnafu)?;
269 Pin::new(&mut stream)
270 .accept()
271 .await
272 .context(HandshakeSnafu)?;
273 Ok(stream)
274 }
275 .boxed(),
276 ),
277 None => StreamState::Accepted(MaybeTlsStream::Raw(stream)),
278 };
279 Self { state, peer_addr }
280 }
281
282 pub async fn handshake(&mut self) -> crate::tls::Result<()> {
284 if let StreamState::Accepting(fut) = &mut self.state {
285 let stream = fut.await?;
286 self.state = StreamState::Accepted(MaybeTlsStream::Tls(stream));
287 }
288
289 Ok(())
290 }
291
292 pub fn set_keepalive(&mut self, keepalive: TcpKeepaliveConfig) -> io::Result<()> {
293 let stream = self.get_ref().ok_or_else(|| {
294 io::Error::new(
295 io::ErrorKind::NotConnected,
296 "Can't set keepalive on connection that has not been accepted yet.",
297 )
298 })?;
299
300 if let Some(time_secs) = keepalive.time_secs {
301 let config =
302 socket2::TcpKeepalive::new().with_time(std::time::Duration::from_secs(time_secs));
303
304 tcp::set_keepalive(stream, &config)?;
305 }
306
307 Ok(())
308 }
309
310 pub fn set_receive_buffer_bytes(&mut self, bytes: usize) -> std::io::Result<()> {
311 let stream = self.get_ref().ok_or_else(|| {
312 io::Error::new(
313 io::ErrorKind::NotConnected,
314 "Can't set receive buffer size on connection that has not been accepted yet.",
315 )
316 })?;
317
318 tcp::set_receive_buffer_size(stream, bytes)
319 }
320
321 fn poll_io<T, F>(self: Pin<&mut Self>, cx: &mut Context, poll_fn: F) -> Poll<io::Result<T>>
322 where
323 F: FnOnce(Pin<&mut MaybeTlsStream<TcpStream>>, &mut Context) -> Poll<io::Result<T>>,
324 {
325 let this = self.get_mut();
326 loop {
327 return match &mut this.state {
328 StreamState::Accepted(stream) => poll_fn(Pin::new(stream), cx),
329 StreamState::Accepting(fut) => match std::task::ready!(fut.as_mut().poll(cx)) {
330 Ok(stream) => {
331 this.state = StreamState::Accepted(MaybeTlsStream::Tls(stream));
332 continue;
333 }
334 Err(error) => {
335 let error = io::Error::other(error);
336 this.state = StreamState::AcceptError(error.to_string());
337 Poll::Ready(Err(error))
338 }
339 },
340 StreamState::AcceptError(error) => {
341 Poll::Ready(Err(io::Error::other(error.clone())))
342 }
343 StreamState::Closed => Poll::Ready(Err(io::ErrorKind::BrokenPipe.into())),
344 };
345 }
346 }
347}
348
349impl AsyncRead for MaybeTlsIncomingStream<TcpStream> {
350 fn poll_read(
351 self: Pin<&mut Self>,
352 cx: &mut Context,
353 buf: &mut ReadBuf<'_>,
354 ) -> Poll<io::Result<()>> {
355 self.poll_io(cx, |s, cx| s.poll_read(cx, buf))
356 }
357}
358
359impl AsyncWrite for MaybeTlsIncomingStream<TcpStream> {
360 fn poll_write(self: Pin<&mut Self>, cx: &mut Context, buf: &[u8]) -> Poll<io::Result<usize>> {
361 self.poll_io(cx, |s, cx| s.poll_write(cx, buf))
362 }
363
364 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
365 self.poll_io(cx, AsyncWrite::poll_flush)
366 }
367
368 fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
369 let this = self.get_mut();
370 match &mut this.state {
371 StreamState::Accepted(stream) => match Pin::new(stream).poll_shutdown(cx) {
372 Poll::Ready(Ok(())) => {
373 this.state = StreamState::Closed;
374 Poll::Ready(Ok(()))
375 }
376 poll_result => poll_result,
377 },
378 StreamState::Accepting(fut) => match std::task::ready!(fut.as_mut().poll(cx)) {
379 Ok(stream) => {
380 this.state = StreamState::Accepted(MaybeTlsStream::Tls(stream));
381 Poll::Pending
382 }
383 Err(error) => {
384 let error = io::Error::other(error);
385 this.state = StreamState::AcceptError(error.to_string());
386 Poll::Ready(Err(error))
387 }
388 },
389 StreamState::AcceptError(error) => Poll::Ready(Err(io::Error::other(error.clone()))),
390 StreamState::Closed => Poll::Ready(Ok(())),
391 }
392 }
393}
394
395#[derive(Debug)]
396pub struct CertificateMetadata {
397 pub country_name: Option<String>,
398 pub state_or_province_name: Option<String>,
399 pub locality_name: Option<String>,
400 pub organization_name: Option<String>,
401 pub organizational_unit_name: Option<String>,
402 pub common_name: Option<String>,
403}
404
405impl CertificateMetadata {
406 pub fn subject(&self) -> String {
407 let mut components = Vec::<String>::with_capacity(6);
408 if let Some(cn) = &self.common_name {
409 components.push(format!("CN={cn}"));
410 }
411 if let Some(ou) = &self.organizational_unit_name {
412 components.push(format!("OU={ou}"));
413 }
414 if let Some(o) = &self.organization_name {
415 components.push(format!("O={o}"));
416 }
417 if let Some(l) = &self.locality_name {
418 components.push(format!("L={l}"));
419 }
420 if let Some(st) = &self.state_or_province_name {
421 components.push(format!("ST={st}"));
422 }
423 if let Some(c) = &self.country_name {
424 components.push(format!("C={c}"));
425 }
426 components.join(",")
427 }
428}
429
430impl From<X509> for CertificateMetadata {
431 fn from(cert: X509) -> Self {
432 let mut subject_metadata: HashMap<String, String> = HashMap::new();
433 for entry in cert.subject_name().entries() {
434 let data_string = entry.data().to_string().unwrap_or_default();
435 subject_metadata.insert(entry.object().to_string(), data_string);
436 }
437 Self {
438 country_name: subject_metadata.get("countryName").cloned(),
439 state_or_province_name: subject_metadata.get("stateOrProvinceName").cloned(),
440 locality_name: subject_metadata.get("localityName").cloned(),
441 organization_name: subject_metadata.get("organizationName").cloned(),
442 organizational_unit_name: subject_metadata.get("organizationalUnitName").cloned(),
443 common_name: subject_metadata.get("commonName").cloned(),
444 }
445 }
446}
447
448#[derive(Clone)]
449pub struct MaybeTlsConnectInfo {
450 pub remote_addr: SocketAddr,
451 pub peer_certs: Option<Vec<Certificate>>,
452}
453
454impl Connected for MaybeTlsIncomingStream<TcpStream> {
455 type ConnectInfo = MaybeTlsConnectInfo;
456
457 fn connect_info(&self) -> Self::ConnectInfo {
458 MaybeTlsConnectInfo {
459 remote_addr: self.peer_addr(),
460 peer_certs: self
461 .ssl_stream()
462 .and_then(|s| s.ssl().peer_cert_chain())
463 .map(|s| {
464 s.into_iter()
465 .filter_map(|c| c.to_pem().ok())
466 .map(Certificate::from_pem)
467 .collect()
468 }),
469 }
470 }
471}
472
473#[cfg(test)]
474mod test {
475 use super::*;
476
477 #[test]
478 fn certificate_metadata_full() {
479 let example_meta = CertificateMetadata {
480 common_name: Some("common".to_owned()),
481 country_name: Some("country".to_owned()),
482 locality_name: Some("locality".to_owned()),
483 organization_name: Some("organization".to_owned()),
484 organizational_unit_name: Some("org_unit".to_owned()),
485 state_or_province_name: Some("state".to_owned()),
486 };
487
488 let expected = format!(
489 "CN={},OU={},O={},L={},ST={},C={}",
490 example_meta.common_name.as_ref().unwrap(),
491 example_meta.organizational_unit_name.as_ref().unwrap(),
492 example_meta.organization_name.as_ref().unwrap(),
493 example_meta.locality_name.as_ref().unwrap(),
494 example_meta.state_or_province_name.as_ref().unwrap(),
495 example_meta.country_name.as_ref().unwrap()
496 );
497 assert_eq!(expected, example_meta.subject());
498 }
499
500 #[test]
501 fn certificate_metadata_partial() {
502 let example_meta = CertificateMetadata {
503 common_name: Some("common".to_owned()),
504 country_name: Some("country".to_owned()),
505 locality_name: None,
506 organization_name: Some("organization".to_owned()),
507 organizational_unit_name: Some("org_unit".to_owned()),
508 state_or_province_name: None,
509 };
510
511 let expected = format!(
512 "CN={},OU={},O={},C={}",
513 example_meta.common_name.as_ref().unwrap(),
514 example_meta.organizational_unit_name.as_ref().unwrap(),
515 example_meta.organization_name.as_ref().unwrap(),
516 example_meta.country_name.as_ref().unwrap()
517 );
518 assert_eq!(expected, example_meta.subject());
519 }
520}