use std::{sync::Arc, time::Duration};
use rustls::client::{WebPkiServerVerifier, danger::ServerCertVerifier};
use crate::dquic::{
log::{QLog, handy::NoopLogger},
param::{ClientParameters, handy::client_parameters},
stream::{ProductStreamsConcurrencyController, handy::ConsistentConcurrency},
token::{TokenSink, handy::NoopTokenRegistry},
};
#[derive(Clone, Default)]
pub enum ServerCertVerifierChoice {
#[default]
Dangerous,
WebPki(Arc<WebPkiServerVerifier>),
Custom(Arc<dyn ServerCertVerifier>),
}
impl std::fmt::Debug for ServerCertVerifierChoice {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Dangerous => f.debug_tuple("Dangerous").finish(),
Self::WebPki(_) => f.debug_tuple("WebPki").finish(),
Self::Custom(_) => f.debug_tuple("Custom").finish(),
}
}
}
impl PartialEq for ServerCertVerifierChoice {
fn eq(&self, other: &Self) -> bool {
match (self, other) {
(Self::Dangerous, Self::Dangerous) => true,
(Self::WebPki(a), Self::WebPki(b)) => Arc::ptr_eq(a, b),
(Self::Custom(a), Self::Custom(b)) => Arc::ptr_eq(a, b),
_ => false,
}
}
}
#[derive(Clone)]
pub struct ClientQuicConfig {
pub defer_idle_timeout: Duration,
pub stream_strategy_factory: Arc<dyn ProductStreamsConcurrencyController>,
pub qlogger: Arc<dyn QLog + Send + Sync>,
pub enable_0rtt: bool,
pub enable_sslkeylog: bool,
pub parameters: ClientParameters,
pub alpns: Vec<Vec<u8>>,
pub token_sink: Arc<dyn TokenSink>,
pub verifier: ServerCertVerifierChoice,
}
impl Default for ClientQuicConfig {
fn default() -> Self {
Self {
defer_idle_timeout: Duration::ZERO,
stream_strategy_factory: Arc::new(ConsistentConcurrency::new),
qlogger: Arc::new(NoopLogger),
enable_0rtt: false,
enable_sslkeylog: false,
parameters: client_parameters(),
alpns: Vec::new(),
token_sink: Arc::new(NoopTokenRegistry),
verifier: ServerCertVerifierChoice::default(),
}
}
}
impl std::fmt::Debug for ClientQuicConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ClientQuicConfig")
.field("defer_idle_timeout", &self.defer_idle_timeout)
.field("enable_0rtt", &self.enable_0rtt)
.field("enable_sslkeylog", &self.enable_sslkeylog)
.field("alpns", &self.alpns.len())
.field("verifier", &self.verifier)
.finish_non_exhaustive()
}
}
impl PartialEq for ClientQuicConfig {
fn eq(&self, other: &Self) -> bool {
self.defer_idle_timeout == other.defer_idle_timeout
&& self.enable_0rtt == other.enable_0rtt
&& self.enable_sslkeylog == other.enable_sslkeylog
&& self.parameters == other.parameters
&& self.alpns == other.alpns
&& self.verifier == other.verifier
&& Arc::ptr_eq(
&self.stream_strategy_factory,
&other.stream_strategy_factory,
)
&& Arc::ptr_eq(&self.qlogger, &other.qlogger)
&& Arc::ptr_eq(&self.token_sink, &other.token_sink)
}
}
#[cfg(test)]
mod client_tests {
use std::{sync::Arc, time::Duration};
use rustls::{
RootCertStore,
client::{WebPkiServerVerifier, danger::ServerCertVerifier},
};
use crate::{
dquic::{client::*, common::*, prelude::handy::ToCertificate},
util::tls::DangerousServerCertVerifier,
};
const CA_CERT: &[u8] = include_bytes!("../../tests/keychain/localhost/ca.cert");
fn root_store_with_ca() -> RootCertStore {
let mut store = RootCertStore::empty();
store.add_parsable_certificates(CA_CERT.to_certificate());
store
}
#[test]
fn test_common_quic_config_default() {
let cfg = CommonQuicConfig::default();
assert_eq!(cfg.defer_idle_timeout, Duration::ZERO);
assert!(!cfg.enable_0rtt);
assert!(!cfg.enable_sslkeylog);
}
#[test]
fn test_common_quic_config_partial_eq_different_timeout() {
let a = CommonQuicConfig::default();
let mut b = a.clone();
b.defer_idle_timeout = Duration::from_secs(30);
assert_ne!(a, b);
}
#[test]
fn test_common_quic_config_clone() {
let a = CommonQuicConfig::default();
let b = a.clone();
assert!(Arc::ptr_eq(
&a.stream_strategy_factory,
&b.stream_strategy_factory
));
assert!(Arc::ptr_eq(&a.qlogger, &b.qlogger));
assert_eq!(a.defer_idle_timeout, b.defer_idle_timeout);
assert_eq!(a.enable_0rtt, b.enable_0rtt);
assert_eq!(a.enable_sslkeylog, b.enable_sslkeylog);
}
#[test]
fn test_verifier_choice_dangerous_ne_webpki() {
let store = root_store_with_ca();
let webpki = WebPkiServerVerifier::builder(Arc::new(store))
.build()
.unwrap();
assert_ne!(
ServerCertVerifierChoice::Dangerous,
ServerCertVerifierChoice::WebPki(webpki)
);
}
#[test]
fn test_verifier_choice_webpki_same_arc_eq() {
let store = root_store_with_ca();
let webpki = WebPkiServerVerifier::builder(Arc::new(store))
.build()
.unwrap();
let a = ServerCertVerifierChoice::WebPki(webpki.clone());
let b = ServerCertVerifierChoice::WebPki(webpki.clone());
assert_eq!(a, b);
}
#[test]
fn test_verifier_choice_webpki_different_arc_ne() {
let store1 = root_store_with_ca();
let store2 = root_store_with_ca();
let webpki1 = WebPkiServerVerifier::builder(Arc::new(store1))
.build()
.unwrap();
let webpki2 = WebPkiServerVerifier::builder(Arc::new(store2))
.build()
.unwrap();
assert_ne!(
ServerCertVerifierChoice::WebPki(webpki1),
ServerCertVerifierChoice::WebPki(webpki2)
);
}
#[test]
fn test_verifier_choice_custom_same_arc_eq() {
let verifier: Arc<dyn ServerCertVerifier> = Arc::new(DangerousServerCertVerifier);
let a = ServerCertVerifierChoice::Custom(verifier.clone());
let b = ServerCertVerifierChoice::Custom(verifier.clone());
assert_eq!(a, b);
}
#[test]
fn test_verifier_choice_custom_different_arc_ne() {
let a: Arc<dyn ServerCertVerifier> = Arc::new(DangerousServerCertVerifier);
let b: Arc<dyn ServerCertVerifier> = Arc::new(DangerousServerCertVerifier);
assert_ne!(
ServerCertVerifierChoice::Custom(a),
ServerCertVerifierChoice::Custom(b)
);
}
#[test]
fn test_verifier_choice_cross_variant_not_equal() {
let verifier: Arc<dyn ServerCertVerifier> = Arc::new(DangerousServerCertVerifier);
assert_ne!(
ServerCertVerifierChoice::Dangerous,
ServerCertVerifierChoice::Custom(verifier.clone())
);
assert_ne!(
ServerCertVerifierChoice::WebPki(
WebPkiServerVerifier::builder(Arc::new(root_store_with_ca()))
.build()
.unwrap()
),
ServerCertVerifierChoice::Custom(verifier)
);
}
#[test]
fn test_verifier_choice_debug_variants() {
let store = root_store_with_ca();
let webpki = WebPkiServerVerifier::builder(Arc::new(store))
.build()
.unwrap();
let dangerous = ServerCertVerifierChoice::Dangerous;
let custom: Arc<dyn ServerCertVerifier> = Arc::new(DangerousServerCertVerifier);
assert_eq!(format!("{:?}", dangerous), "Dangerous");
assert_eq!(
format!("{:?}", ServerCertVerifierChoice::WebPki(webpki)),
"WebPki"
);
assert_eq!(
format!("{:?}", ServerCertVerifierChoice::Custom(custom)),
"Custom"
);
}
#[test]
fn test_verifier_choice_default_is_dangerous() {
assert_eq!(
ServerCertVerifierChoice::default(),
ServerCertVerifierChoice::Dangerous
);
}
#[test]
fn test_client_quic_config_default() {
let cfg = ClientQuicConfig::default();
assert_eq!(cfg.defer_idle_timeout, Duration::ZERO);
assert!(!cfg.enable_0rtt);
assert!(!cfg.enable_sslkeylog);
assert!(
matches!(&cfg.verifier, ServerCertVerifierChoice::Dangerous),
"default verifier should be Dangerous"
);
assert!(cfg.alpns.is_empty(), "default alpns should be empty");
}
#[test]
fn test_client_quic_config_partial_eq_different_timeout() {
let a = ClientQuicConfig::default();
let mut b = a.clone();
b.defer_idle_timeout = Duration::from_secs(99);
assert_ne!(a, b);
}
#[test]
fn test_client_quic_config_partial_eq_different_verifier() {
let a = ClientQuicConfig::default();
let store = root_store_with_ca();
let webpki = WebPkiServerVerifier::builder(Arc::new(store))
.build()
.unwrap();
let mut custom = a.clone();
custom.verifier = ServerCertVerifierChoice::Custom(Arc::new(DangerousServerCertVerifier));
assert_ne!(a, custom);
let mut webpki_choice = a.clone();
webpki_choice.verifier = ServerCertVerifierChoice::WebPki(webpki);
assert_ne!(a, webpki_choice);
}
#[test]
fn test_client_quic_config_partial_eq_different_components() {
let a = ClientQuicConfig::default();
let mut strategy = a.clone();
strategy.stream_strategy_factory = Arc::new(ConsistentConcurrency::new);
assert_ne!(a, strategy);
let mut qlogger = a.clone();
qlogger.qlogger = Arc::new(NoopLogger);
assert_ne!(a, qlogger);
let mut token_sink = a.clone();
token_sink.token_sink = Arc::new(NoopTokenRegistry);
assert_ne!(a, token_sink);
}
#[test]
fn test_client_quic_config_debug() {
let cfg = ClientQuicConfig {
alpns: vec![b"h3".to_vec()],
..ClientQuicConfig::default()
};
let rendered = format!("{cfg:?}");
assert!(rendered.contains("ClientQuicConfig"));
assert!(rendered.contains("defer_idle_timeout: 0ns"));
assert!(rendered.contains("enable_0rtt: false"));
assert!(rendered.contains("enable_sslkeylog: false"));
assert!(rendered.contains("alpns: 1"));
assert!(rendered.contains("verifier: Dangerous"));
assert!(rendered.contains(".."));
assert!(!rendered.contains("stream_strategy_factory"));
}
#[test]
fn test_client_quic_config_clone() {
let a = ClientQuicConfig::default();
let b = a.clone();
assert!(Arc::ptr_eq(
&a.stream_strategy_factory,
&b.stream_strategy_factory
));
assert!(Arc::ptr_eq(&a.qlogger, &b.qlogger));
assert!(Arc::ptr_eq(&a.token_sink, &b.token_sink));
assert_eq!(a.defer_idle_timeout, b.defer_idle_timeout);
assert_eq!(a.enable_0rtt, b.enable_0rtt);
assert_eq!(a.enable_sslkeylog, b.enable_sslkeylog);
assert_eq!(a.parameters, b.parameters);
assert_eq!(a.alpns, b.alpns);
assert_eq!(a.verifier, b.verifier);
}
#[test]
fn test_client_quic_config_mutate_does_not_affect_clone() {
let a = ClientQuicConfig::default();
let mut b = a.clone();
b.defer_idle_timeout = Duration::from_secs(99);
b.alpns.push(b"h3".to_vec());
assert_eq!(a.defer_idle_timeout, Duration::ZERO);
assert!(a.alpns.is_empty());
assert_eq!(b.defer_idle_timeout, Duration::from_secs(99));
assert!(!b.alpns.is_empty());
}
#[test]
fn test_client_quic_config_mutate_arc_fields_does_not_affect_clone() {
let a = ClientQuicConfig::default();
let mut b = a.clone();
b.stream_strategy_factory = Arc::new(ConsistentConcurrency::new);
b.qlogger = Arc::new(NoopLogger);
b.token_sink = Arc::new(NoopTokenRegistry);
assert_eq!(a.defer_idle_timeout, b.defer_idle_timeout);
assert_eq!(a.enable_0rtt, b.enable_0rtt);
assert_eq!(a.enable_sslkeylog, b.enable_sslkeylog);
assert_eq!(a.parameters, b.parameters);
assert_eq!(a.alpns, b.alpns);
assert_eq!(a.verifier, b.verifier);
assert!(!Arc::ptr_eq(
&a.stream_strategy_factory,
&b.stream_strategy_factory
));
assert!(!Arc::ptr_eq(&a.qlogger, &b.qlogger));
assert!(!Arc::ptr_eq(&a.token_sink, &b.token_sink));
}
}