use super::HasConnectionInfo;
#[cfg(all(feature = "server", feature = "tls"))]
pub use self::channel::{TlsConnectionInfoReceiver, TlsConnectionInfoSender, channel};
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct TlsConnectionInfo {
pub server_name: Option<String>,
pub validated_server_name: bool,
pub alpn: Option<String>,
}
#[cfg(feature = "tls")]
impl TlsConnectionInfo {
pub fn server(server_info: &rustls::ServerConnection) -> Self {
let server_name = server_info
.server_name()
.map(|s| s.to_string())
.filter(|s| !s.is_empty());
let alpn = server_info
.alpn_protocol()
.and_then(|s| std::str::from_utf8(s).ok())
.and_then(|s| s.parse().ok());
Self {
server_name,
validated_server_name: false,
alpn,
}
}
pub fn client(client_info: &rustls::ClientConnection) -> Self {
let alpn = client_info
.alpn_protocol()
.and_then(|s| std::str::from_utf8(s).ok())
.and_then(|s| s.parse().ok());
Self {
server_name: None,
validated_server_name: false,
alpn,
}
}
#[cfg(test)]
#[allow(dead_code)]
pub(crate) fn new_server(
server_name: Option<String>,
validated_server_name: bool,
alpn: Option<String>,
) -> Self {
Self {
server_name,
validated_server_name,
alpn,
}
}
pub fn validated(&mut self) {
self.validated_server_name = true;
}
}
pub trait HasTlsConnectionInfo: HasConnectionInfo {
fn tls_info(&self) -> Option<&TlsConnectionInfo>;
}
#[cfg(all(feature = "server", feature = "tls"))]
mod channel {
use std::{
ops::{Deref, DerefMut},
sync::Arc,
};
use tokio::sync::RwLock;
use crate::info::TlsConnectionInfo;
pub fn channel() -> (TlsConnectionInfoSender, TlsConnectionInfoReceiver) {
let (tx, rx) = tokio::sync::oneshot::channel();
(
TlsConnectionInfoSender { tx: Some(tx) },
TlsConnectionInfoReceiver::new(rx),
)
}
#[derive(Debug)]
pub struct TlsConnectionInfoSender {
tx: Option<tokio::sync::oneshot::Sender<TlsConnectionInfo>>,
}
impl TlsConnectionInfoSender {
pub fn send(&mut self, info: TlsConnectionInfo) {
if let Some(tx) = self.tx.take() {
let _ = tx.send(info);
}
}
}
#[derive(Debug)]
enum State {
Pending(tokio::sync::oneshot::Receiver<TlsConnectionInfo>),
Received(TlsConnectionInfo),
Empty,
}
#[derive(Debug, Clone)]
pub struct TlsConnectionInfoReceiver {
state: Arc<RwLock<State>>,
}
impl TlsConnectionInfoReceiver {
pub fn empty() -> Self {
Self {
state: Arc::new(RwLock::new(State::Empty)),
}
}
pub(crate) fn new(inner: tokio::sync::oneshot::Receiver<TlsConnectionInfo>) -> Self {
Self {
state: Arc::new(RwLock::new(State::Pending(inner))),
}
}
pub(crate) async fn recv(&self) -> Option<TlsConnectionInfo> {
{
let state = self.state.read().await;
match state.deref() {
State::Pending(_) => {}
State::Received(info) => return Some(info.clone()),
State::Empty => return None,
};
}
let mut state = self.state.write().await;
let rx = match state.deref_mut() {
State::Pending(rx) => rx,
State::Received(info) => {
return Some(info.clone());
}
State::Empty => {
return None;
}
};
let tls = rx
.await
.expect("connection info was never sent and is not available");
*state = State::Received(tls.clone());
Some(tls)
}
}
}
#[cfg(all(test, feature = "tls"))]
mod tests {
use crate::fixtures;
use std::sync::Arc;
#[test]
fn tls_server_info() {
fixtures::tls_install_default();
#[derive(Debug, Clone)]
struct NoCertResolver;
impl rustls::server::ResolvesServerCert for NoCertResolver {
fn resolve(
&self,
_client_hello: rustls::server::ClientHello,
) -> Option<Arc<rustls::sign::CertifiedKey>> {
None
}
}
let cfg = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(Arc::new(NoCertResolver));
let server_info = rustls::ServerConnection::new(Arc::new(cfg)).unwrap();
let info = super::TlsConnectionInfo::server(&server_info);
assert_eq!(info.server_name, None);
assert!(!info.validated_server_name);
assert_eq!(info.alpn, None);
}
#[test]
fn tls_client_info() {
fixtures::tls_install_default();
let root_store = rustls::RootCertStore {
roots: webpki_roots::TLS_SERVER_ROOTS.into(),
};
let cfg = rustls::ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
let client_info =
rustls::ClientConnection::new(Arc::new(cfg), "example.com".try_into().unwrap())
.unwrap();
let info = super::TlsConnectionInfo::client(&client_info);
assert_eq!(info.server_name, None);
assert!(!info.validated_server_name);
assert_eq!(info.alpn, None);
}
}