use std::collections::hash_map::DefaultHasher;
use std::fmt;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
mod cert;
pub use cert::{Certificate, PemItem, PrivateKey, parse_pem};
#[cfg(feature = "_rustls")]
pub(crate) mod rustls;
#[cfg(feature = "native-tls")]
pub(crate) mod native_tls;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
#[non_exhaustive]
pub enum TlsProvider {
#[default]
Rustls,
NativeTls,
}
impl TlsProvider {
pub(crate) fn is_feature_enabled(&self) -> bool {
match self {
TlsProvider::Rustls => {
cfg!(feature = "_rustls")
}
TlsProvider::NativeTls => {
cfg!(feature = "native-tls")
}
}
}
pub(crate) fn feature_name(&self) -> &'static str {
match self {
TlsProvider::Rustls => "rustls",
TlsProvider::NativeTls => "native-tls",
}
}
}
#[derive(Clone)]
pub struct TlsConfig {
provider: TlsProvider,
client_cert: Option<ClientCert>,
root_certs: RootCerts,
use_sni: bool,
disable_verification: bool,
#[cfg(feature = "_rustls")]
rustls_crypto_provider: Option<Arc<::rustls::crypto::CryptoProvider>>,
}
impl TlsConfig {
pub(crate) fn can_share_pool_with(&self, agent: &Self) -> bool {
let Self {
provider,
client_cert,
root_certs,
use_sni,
disable_verification,
#[cfg(feature = "_rustls")]
rustls_crypto_provider,
} = self;
let same_client = match (client_cert, &agent.client_cert) {
(None, None) => true,
(Some(a), Some(b)) => Arc::ptr_eq(&a.0, &b.0),
_ => false,
};
let same_roots = match (root_certs, &agent.root_certs) {
(RootCerts::WebPki, RootCerts::WebPki)
| (RootCerts::PlatformVerifier, RootCerts::PlatformVerifier) => true,
(RootCerts::Specific(a), RootCerts::Specific(b)) => Arc::ptr_eq(a, b),
_ => false,
};
#[cfg(feature = "_rustls")]
match (rustls_crypto_provider, &agent.rustls_crypto_provider) {
(None, None) => {}
(Some(a), Some(b)) if Arc::ptr_eq(a, b) => {}
_ => return false,
}
*provider == agent.provider
&& same_client
&& same_roots
&& *use_sni == agent.use_sni
&& *disable_verification == agent.disable_verification
}
pub fn builder() -> TlsConfigBuilder {
TlsConfigBuilder {
config: TlsConfig::default(),
}
}
pub(crate) fn hash_value(&self) -> u64 {
let mut hasher = DefaultHasher::new();
self.hash(&mut hasher);
hasher.finish()
}
}
impl TlsConfig {
pub fn provider(&self) -> TlsProvider {
self.provider
}
pub fn client_cert(&self) -> Option<&ClientCert> {
self.client_cert.as_ref()
}
pub fn root_certs(&self) -> &RootCerts {
&self.root_certs
}
pub fn use_sni(&self) -> bool {
self.use_sni
}
pub fn disable_verification(&self) -> bool {
self.disable_verification
}
#[cfg(feature = "_rustls")]
pub fn unversioned_rustls_crypto_provider(
&self,
) -> &Option<Arc<::rustls::crypto::CryptoProvider>> {
&self.rustls_crypto_provider
}
}
pub struct TlsConfigBuilder {
config: TlsConfig,
}
impl TlsConfigBuilder {
pub fn provider(mut self, v: TlsProvider) -> Self {
self.config.provider = v;
self
}
pub fn client_cert(mut self, v: Option<ClientCert>) -> Self {
self.config.client_cert = v;
self
}
pub fn root_certs(mut self, v: RootCerts) -> Self {
self.config.root_certs = v;
self
}
pub fn use_sni(mut self, v: bool) -> Self {
self.config.use_sni = v;
self
}
pub fn disable_verification(mut self, v: bool) -> Self {
self.config.disable_verification = v;
self
}
#[cfg(feature = "_rustls")]
pub fn unversioned_rustls_crypto_provider(
mut self,
v: Arc<::rustls::crypto::CryptoProvider>,
) -> Self {
self.config.rustls_crypto_provider = Some(v);
self
}
pub fn build(self) -> TlsConfig {
self.config
}
}
#[derive(Debug, Clone, Hash)]
pub struct ClientCert(Arc<(Vec<Certificate<'static>>, PrivateKey<'static>)>);
impl ClientCert {
pub fn new_with_certs(chain: &[Certificate<'static>], key: PrivateKey<'static>) -> Self {
Self(Arc::new((chain.to_vec(), key)))
}
pub fn certs(&self) -> &[Certificate<'static>] {
&self.0.0
}
pub fn private_key(&self) -> &PrivateKey<'static> {
&self.0.1
}
}
#[derive(Debug, Clone, Hash)]
#[non_exhaustive]
pub enum RootCerts {
Specific(Arc<Vec<Certificate<'static>>>),
PlatformVerifier,
WebPki,
}
impl RootCerts {
pub fn new_with_certs(certs: &[Certificate<'static>]) -> Self {
certs.iter().cloned().into()
}
}
impl<I: IntoIterator<Item = Certificate<'static>>> From<I> for RootCerts {
fn from(value: I) -> Self {
RootCerts::Specific(Arc::new(value.into_iter().collect()))
}
}
impl Default for TlsConfig {
fn default() -> Self {
let provider = TlsProvider::default();
Self {
provider,
client_cert: None,
root_certs: RootCerts::WebPki,
use_sni: true,
disable_verification: false,
#[cfg(feature = "_rustls")]
rustls_crypto_provider: None,
}
}
}
impl fmt::Debug for TlsConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TlsConfig")
.field("provider", &self.provider)
.field("client_cert", &self.client_cert)
.field("root_certs", &self.root_certs)
.field("use_sni", &self.use_sni)
.field("disable_verification", &self.disable_verification)
.finish()
}
}
impl Hash for TlsConfig {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.provider.hash(state);
self.client_cert.hash(state);
self.root_certs.hash(state);
self.use_sni.hash(state);
self.disable_verification.hash(state);
#[cfg(feature = "_rustls")]
if let Some(arc) = &self.rustls_crypto_provider {
(Arc::as_ptr(arc) as usize).hash(state);
}
}
}
#[cfg(test)]
mod test {
use super::*;
use assert_no_alloc::*;
#[test]
fn specific_roots_pool_by_identity() {
let roots = || RootCerts::new_with_certs(&[Certificate::from_der(b"root")]);
let a = TlsConfig::builder().root_certs(roots()).build();
assert!(a.can_share_pool_with(&a.clone()));
let b = TlsConfig::builder().root_certs(roots()).build();
assert!(!a.can_share_pool_with(&b));
}
#[test]
#[cfg(feature = "_rustls")]
fn crypto_providers_pool_by_identity() {
let provider = || Arc::new(::rustls::crypto::aws_lc_rs::default_provider());
let a = TlsConfig::builder()
.unversioned_rustls_crypto_provider(provider())
.build();
assert!(a.can_share_pool_with(&a.clone()));
let b = TlsConfig::builder()
.unversioned_rustls_crypto_provider(provider())
.build();
assert!(!a.can_share_pool_with(&b));
assert!(!a.can_share_pool_with(&TlsConfig::default()));
assert!(!TlsConfig::default().can_share_pool_with(&a));
}
#[test]
fn tls_config_clone_does_not_allocate() {
let c = TlsConfig::default();
assert_no_alloc(|| c.clone());
}
#[cfg(any(feature = "_rustls", feature = "native-tls"))]
mod handshake {
use std::sync::Arc;
use crate::tls::{TlsConfig, TlsProvider};
use crate::transport::time::{Duration, Instant};
use crate::transport::{Buffers, ConnectionDetails, Connector, LazyBuffers};
use crate::transport::{NextTimeout, Transport};
use crate::unversioned::resolver::{DefaultResolver, Resolver};
use crate::{Agent, Error, Timeout};
#[derive(Debug)]
struct TimeoutTransport {
buffers: LazyBuffers,
timeout: NextTimeout,
}
impl Transport for TimeoutTransport {
fn buffers(&mut self) -> &mut dyn Buffers {
&mut self.buffers
}
fn transmit_output(
&mut self,
_amount: usize,
timeout: NextTimeout,
) -> Result<(), Error> {
assert_eq!(timeout, self.timeout);
Err(Error::Timeout(timeout.reason))
}
fn await_input(&mut self, timeout: NextTimeout) -> Result<bool, Error> {
assert_eq!(timeout, self.timeout);
Err(Error::Timeout(timeout.reason))
}
fn is_open(&mut self) -> bool {
true
}
}
fn check_timeout(connector: impl Connector<TimeoutTransport>, tls_config: TlsConfig) {
let config = Agent::config_builder()
.proxy(None)
.tls_config(tls_config)
.build();
let uri = "https://example.com/".parse().unwrap();
let resolver = DefaultResolver::default();
for reason in [Timeout::Connect, Timeout::Global] {
let timeout = NextTimeout {
after: Duration::from_secs(7),
reason,
};
let details = ConnectionDetails {
uri: &uri,
addrs: resolver.empty(),
resolver: &resolver,
config: &config,
request_level: false,
now: Instant::now(),
timeout,
current_time: Arc::new(Instant::now),
run_connector: Arc::new(|_| unreachable!("TLS must use the chained transport")),
};
let transport = TimeoutTransport {
buffers: LazyBuffers::new(1024, 1024),
timeout,
};
let error = connector.connect(&details, Some(transport)).unwrap_err();
assert!(
matches!(error, Error::Timeout(actual) if actual == reason),
"{error:?}"
);
}
}
#[test]
#[cfg(feature = "_rustls")]
fn rustls_uses_connection_timeout() {
let tls_config = TlsConfig::builder()
.provider(TlsProvider::Rustls)
.disable_verification(true)
.unversioned_rustls_crypto_provider(Arc::new(
::rustls::crypto::aws_lc_rs::default_provider(),
))
.build();
check_timeout(crate::tls::rustls::RustlsConnector::default(), tls_config);
}
#[test]
#[cfg(feature = "native-tls")]
fn native_tls_uses_connection_timeout() {
let tls_config = TlsConfig::builder()
.provider(TlsProvider::NativeTls)
.disable_verification(true)
.build();
check_timeout(
crate::tls::native_tls::NativeTlsConnector::default(),
tls_config,
);
}
}
}