xref: /tonic/tonic/src/transport/server/conn.rs (revision b28b59b2)
1 use std::net::SocketAddr;
2 use tokio::net::TcpStream;
3 
4 #[cfg(feature = "_tls-any")]
5 use std::sync::Arc;
6 #[cfg(feature = "_tls-any")]
7 use tokio_rustls::rustls::pki_types::CertificateDer;
8 #[cfg(feature = "_tls-any")]
9 use tokio_rustls::server::TlsStream;
10 
11 /// Trait that connected IO resources implement and use to produce info about the connection.
12 ///
13 /// The goal for this trait is to allow users to implement
14 /// custom IO types that can still provide the same connection
15 /// metadata.
16 ///
17 /// # Example
18 ///
19 /// The `ConnectInfo` returned will be accessible through [request extensions][ext]:
20 ///
21 /// ```
22 /// use tonic::{Request, transport::server::Connected};
23 ///
24 /// // A `Stream` that yields connections
25 /// struct MyConnector {}
26 ///
27 /// // Return metadata about the connection as `MyConnectInfo`
28 /// impl Connected for MyConnector {
29 ///     type ConnectInfo = MyConnectInfo;
30 ///
31 ///     fn connect_info(&self) -> Self::ConnectInfo {
32 ///         MyConnectInfo {}
33 ///     }
34 /// }
35 ///
36 /// #[derive(Clone)]
37 /// struct MyConnectInfo {
38 ///     // Metadata about your connection
39 /// }
40 ///
41 /// // The connect info can be accessed through request extensions:
42 /// # fn foo(request: Request<()>) {
43 /// let connect_info: &MyConnectInfo = request
44 ///     .extensions()
45 ///     .get::<MyConnectInfo>()
46 ///     .expect("bug in tonic");
47 /// # }
48 /// ```
49 ///
50 /// [ext]: crate::Request::extensions
51 pub trait Connected {
52     /// The connection info type the IO resources generates.
53     // all these bounds are necessary to set this as a request extension
54     type ConnectInfo: Clone + Send + Sync + 'static;
55 
56     /// Create type holding information about the connection.
connect_info(&self) -> Self::ConnectInfo57     fn connect_info(&self) -> Self::ConnectInfo;
58 }
59 
60 /// Connection info for standard TCP streams.
61 ///
62 /// This type will be accessible through [request extensions][ext] if you're using the default
63 /// non-TLS connector.
64 ///
65 /// See [`Connected`] for more details.
66 ///
67 /// [ext]: crate::Request::extensions
68 #[derive(Debug, Clone)]
69 pub struct TcpConnectInfo {
70     /// Returns the local address of this connection.
71     pub local_addr: Option<SocketAddr>,
72     /// Returns the remote (peer) address of this connection.
73     pub remote_addr: Option<SocketAddr>,
74 }
75 
76 impl TcpConnectInfo {
77     /// Return the local address the IO resource is connected.
local_addr(&self) -> Option<SocketAddr>78     pub fn local_addr(&self) -> Option<SocketAddr> {
79         self.local_addr
80     }
81 
82     /// Return the remote address the IO resource is connected too.
remote_addr(&self) -> Option<SocketAddr>83     pub fn remote_addr(&self) -> Option<SocketAddr> {
84         self.remote_addr
85     }
86 }
87 
88 impl Connected for TcpStream {
89     type ConnectInfo = TcpConnectInfo;
90 
connect_info(&self) -> Self::ConnectInfo91     fn connect_info(&self) -> Self::ConnectInfo {
92         TcpConnectInfo {
93             local_addr: self.local_addr().ok(),
94             remote_addr: self.peer_addr().ok(),
95         }
96     }
97 }
98 
99 impl Connected for tokio::io::DuplexStream {
100     type ConnectInfo = ();
101 
connect_info(&self) -> Self::ConnectInfo102     fn connect_info(&self) -> Self::ConnectInfo {}
103 }
104 
105 #[cfg(feature = "_tls-any")]
106 impl<T> Connected for TlsStream<T>
107 where
108     T: Connected,
109 {
110     type ConnectInfo = TlsConnectInfo<T::ConnectInfo>;
111 
connect_info(&self) -> Self::ConnectInfo112     fn connect_info(&self) -> Self::ConnectInfo {
113         let (inner, session) = self.get_ref();
114         let inner = inner.connect_info();
115 
116         let certs = session
117             .peer_certificates()
118             .map(|certs| certs.to_owned().into());
119 
120         TlsConnectInfo { inner, certs }
121     }
122 }
123 
124 /// Connection info for TLS streams.
125 ///
126 /// This type will be accessible through [request extensions][ext] if you're using a TLS connector.
127 ///
128 /// See [`Connected`] for more details.
129 ///
130 /// [ext]: crate::Request::extensions
131 #[cfg(feature = "_tls-any")]
132 #[derive(Debug, Clone)]
133 pub struct TlsConnectInfo<T> {
134     inner: T,
135     certs: Option<Arc<Vec<CertificateDer<'static>>>>,
136 }
137 
138 #[cfg(feature = "_tls-any")]
139 impl<T> TlsConnectInfo<T> {
140     /// Get a reference to the underlying connection info.
get_ref(&self) -> &T141     pub fn get_ref(&self) -> &T {
142         &self.inner
143     }
144 
145     /// Get a mutable reference to the underlying connection info.
get_mut(&mut self) -> &mut T146     pub fn get_mut(&mut self) -> &mut T {
147         &mut self.inner
148     }
149 
150     /// Return the set of connected peer TLS certificates.
peer_certs(&self) -> Option<Arc<Vec<CertificateDer<'static>>>>151     pub fn peer_certs(&self) -> Option<Arc<Vec<CertificateDer<'static>>>> {
152         self.certs.clone()
153     }
154 }
155