Skip to main content

vector_core/tls/
incoming.rs

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    /// Bind a listener, optionally sharing a [`TlsAcceptorReloader`] so its acceptor can be swapped
62    /// at runtime. When `reloader` is `Some` and TLS is enabled, connections accepted by the
63    /// returned listener use the acceptor the reloader currently holds; otherwise the acceptor is
64    /// fixed for the lifetime of the listener.
65    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    // BoxFuture doesn't allow access to the inner stream, but users
203    // of MaybeTlsIncomingStream want access to the peer address while
204    // still handshaking, so we have to cache it here.
205    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    /// None if connection still hasn't been established.
221    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    // Explicit handshake method
283    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}