#[cfg(test)]
mod config_test;
use crate::cipher_suite::*;
use crate::conn::{
DEFAULT_MAXIMUM_RETRANSMIT_NUMBER, DEFAULT_REPLAY_PROTECTION_WINDOW, INITIAL_TICKER_INTERVAL,
};
use crate::crypto::*;
use crate::curve::named_curve::NamedCurve;
use crate::extension::extension_use_srtp::SrtpProtectionProfile;
use crate::signature_hash_algorithm::{
SignatureHashAlgorithm, SignatureScheme, parse_signature_schemes,
};
use crypto::RTCCryptoProvider;
use log::warn;
use shared::error::*;
use std::collections::HashMap;
use std::fmt;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use rustls::client::danger::ServerCertVerifier;
use rustls::pki_types::CertificateDer;
use rustls::server::danger::ClientCertVerifier;
#[derive(Clone)]
pub struct RustlsVerifierAdapter {
provider: Arc<rustls::crypto::CryptoProvider>,
}
impl RustlsVerifierAdapter {
#[must_use]
pub fn new(provider: Arc<rustls::crypto::CryptoProvider>) -> Self {
Self { provider }
}
#[cfg(feature = "crypto-ring")]
#[must_use]
pub fn ring() -> Self {
Self::new(Arc::new(rustls::crypto::ring::default_provider()))
}
#[cfg(feature = "crypto-aws-lc-rs")]
#[must_use]
pub fn aws_lc_rs() -> Self {
Self::new(Arc::new(rustls::crypto::aws_lc_rs::default_provider()))
}
}
fn default_verifier_adapter() -> Option<RustlsVerifierAdapter> {
#[cfg(feature = "crypto-ring")]
{
Some(RustlsVerifierAdapter::ring())
}
#[cfg(all(not(feature = "crypto-ring"), feature = "crypto-aws-lc-rs"))]
{
Some(RustlsVerifierAdapter::aws_lc_rs())
}
#[cfg(not(any(feature = "crypto-ring", feature = "crypto-aws-lc-rs")))]
{
None
}
}
fn server_cert_verifier(
roots: std::sync::Arc<rustls::RootCertStore>,
adapter: &RustlsVerifierAdapter,
) -> Result<std::sync::Arc<dyn ServerCertVerifier>> {
let verifier = rustls::client::WebPkiServerVerifier::builder_with_provider(
roots,
adapter.provider.clone(),
)
.build()
.map_err(|err| Error::Other(format!("rustls server cert verifier: {err}")))?;
Ok(verifier)
}
#[derive(Clone)]
pub struct ConfigBuilder {
crypto_provider: Option<Arc<dyn RTCCryptoProvider>>,
certificates: Vec<Certificate>,
cipher_suites: Vec<CipherSuiteId>,
signature_schemes: Vec<SignatureScheme>,
srtp_protection_profiles: Vec<SrtpProtectionProfile>,
client_auth: ClientAuthType,
extended_master_secret: ExtendedMasterSecretType,
flight_interval: Duration,
psk: Option<PskCallback>,
psk_identity_hint: Option<Vec<u8>>,
insecure_skip_verify: bool,
insecure_hashes: bool,
insecure_verification: bool,
verify_peer_certificate: Option<VerifyPeerCertificateFn>,
roots_cas: rustls::RootCertStore,
client_cas: rustls::RootCertStore,
verifier_adapter: Option<RustlsVerifierAdapter>,
server_name: String,
mtu: usize,
replay_protection_window: usize,
}
impl Default for ConfigBuilder {
fn default() -> Self {
Self {
crypto_provider: None,
certificates: vec![],
cipher_suites: vec![],
signature_schemes: vec![],
srtp_protection_profiles: vec![],
client_auth: ClientAuthType::default(),
extended_master_secret: ExtendedMasterSecretType::default(),
flight_interval: Duration::default(),
psk: None,
psk_identity_hint: None,
insecure_skip_verify: false,
insecure_hashes: false,
insecure_verification: false,
verify_peer_certificate: None,
roots_cas: rustls::RootCertStore::empty(),
client_cas: rustls::RootCertStore::empty(),
verifier_adapter: default_verifier_adapter(),
server_name: String::default(),
mtu: 0,
replay_protection_window: 0,
}
}
}
impl ConfigBuilder {
pub fn with_crypto_provider(mut self, provider: Arc<dyn RTCCryptoProvider>) -> Self {
self.crypto_provider = Some(provider);
self
}
pub fn with_certificates(mut self, certificates: Vec<Certificate>) -> Self {
self.certificates = certificates;
self
}
pub fn with_cipher_suites(mut self, cipher_suites: Vec<CipherSuiteId>) -> Self {
self.cipher_suites = cipher_suites;
self
}
pub fn with_signature_schemes(mut self, signature_schemes: Vec<SignatureScheme>) -> Self {
self.signature_schemes = signature_schemes;
self
}
pub fn with_srtp_protection_profiles(
mut self,
srtp_protection_profiles: Vec<SrtpProtectionProfile>,
) -> Self {
self.srtp_protection_profiles = srtp_protection_profiles;
self
}
pub fn with_client_auth(mut self, client_auth: ClientAuthType) -> Self {
self.client_auth = client_auth;
self
}
pub fn with_extended_master_secret(
mut self,
extended_master_secret: ExtendedMasterSecretType,
) -> Self {
self.extended_master_secret = extended_master_secret;
self
}
pub fn with_flight_interval(mut self, flight_interval: Duration) -> Self {
self.flight_interval = flight_interval;
self
}
pub fn with_psk(mut self, psk: Option<PskCallback>) -> Self {
self.psk = psk;
self
}
pub fn with_psk_identity_hint(mut self, psk_identity_hint: Option<Vec<u8>>) -> Self {
self.psk_identity_hint = psk_identity_hint;
self
}
pub fn with_insecure_skip_verify(mut self, insecure_skip_verify: bool) -> Self {
self.insecure_skip_verify = insecure_skip_verify;
self
}
pub fn with_insecure_hashes(mut self, insecure_hashes: bool) -> Self {
self.insecure_hashes = insecure_hashes;
self
}
pub fn with_insecure_verification(mut self, insecure_verification: bool) -> Self {
self.insecure_verification = insecure_verification;
self
}
pub fn with_verify_peer_certificate(
mut self,
verify_peer_certificate: Option<VerifyPeerCertificateFn>,
) -> Self {
self.verify_peer_certificate = verify_peer_certificate;
self
}
pub fn with_roots_cas(mut self, roots_cas: rustls::RootCertStore) -> Self {
self.roots_cas = roots_cas;
self
}
pub fn with_client_cas(mut self, client_cas: rustls::RootCertStore) -> Self {
self.client_cas = client_cas;
self
}
pub fn with_rustls_verifier_adapter(mut self, adapter: RustlsVerifierAdapter) -> Self {
self.verifier_adapter = Some(adapter);
self
}
pub fn with_server_name(mut self, server_name: String) -> Self {
self.server_name = server_name;
self
}
pub fn with_mtu(mut self, mtu: usize) -> Self {
self.mtu = mtu;
self
}
pub fn with_replay_protection_window(mut self, replay_protection_window: usize) -> Self {
self.replay_protection_window = replay_protection_window;
self
}
}
pub(crate) const DEFAULT_MTU: usize = 1200;
pub(crate) type PskCallback = Arc<dyn (Fn(&[u8]) -> Result<Vec<u8>>) + Send + Sync>;
#[derive(Debug, Default, Copy, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum ClientAuthType {
#[default]
NoClientCert = 0,
RequestClientCert = 1,
RequireAnyClientCert = 2,
VerifyClientCertIfGiven = 3,
RequireAndVerifyClientCert = 4,
}
#[derive(Debug, Default, PartialEq, Eq, Copy, Clone)]
pub enum ExtendedMasterSecretType {
#[default]
Request = 0,
Require = 1,
Disable = 2,
}
impl ConfigBuilder {
fn validate(&self, is_client: bool) -> Result<()> {
if is_client && self.psk.is_some() && self.psk_identity_hint.is_none() {
return Err(Error::ErrPskAndIdentityMustBeSetForClient);
}
if !is_client && self.psk.is_none() && self.certificates.is_empty() {
return Err(Error::ErrServerMustHaveCertificate);
}
if !self.certificates.is_empty() && self.psk.is_some() {
return Err(Error::ErrPskAndCertificate);
}
if self.psk_identity_hint.is_some() && self.psk.is_none() {
return Err(Error::ErrIdentityNoPsk);
}
parse_cipher_suites(&self.cipher_suites, self.psk.is_none(), self.psk.is_some())?;
Ok(())
}
pub fn build(
mut self,
is_client: bool,
remote_addr: Option<SocketAddr>,
) -> Result<HandshakeConfig> {
let crypto_provider = self.crypto_provider.take().ok_or_else(|| {
Error::Crypto(
"no crypto provider configured: call ConfigBuilder::with_crypto_provider"
.to_owned(),
)
})?;
self.validate(is_client)?;
let mut local_cipher_suites: Vec<CipherSuiteId> =
parse_cipher_suites(&self.cipher_suites, self.psk.is_none(), self.psk.is_some())?
.iter()
.map(|cs| cs.id())
.filter(|id| id.supported_by(crypto_provider.crypto()))
.collect();
if local_cipher_suites.is_empty() {
return Err(Error::ErrNoAvailableCipherSuites);
}
let sigs: Vec<u16> = self.signature_schemes.iter().map(|x| *x as u16).collect();
let local_signature_schemes: Vec<_> = parse_signature_schemes(&sigs, self.insecure_hashes)?
.into_iter()
.filter(|algorithm| {
algorithm.crypto_scheme().is_ok_and(|scheme| {
crypto_provider
.crypto()
.supports(crypto::CryptoAlgorithm::Signature(scheme))
})
})
.collect();
if self.psk.is_none() && local_signature_schemes.is_empty() {
return Err(Error::ErrNoAvailableSignatureSchemes);
}
let local_named_curves: Vec<_> = [NamedCurve::P256, NamedCurve::X25519, NamedCurve::P384]
.into_iter()
.filter(|curve| {
curve.crypto_algorithm().is_ok_and(|algorithm| {
crypto_provider
.crypto()
.supports(crypto::CryptoAlgorithm::KeyExchange(algorithm))
})
})
.collect();
if self.psk.is_none() && local_named_curves.is_empty() {
return Err(Error::ErrNoAvailableCipherSuites);
}
if !is_client && self.psk.is_none() {
let signing_key = &self.certificates[0].private_key.signing_key;
local_cipher_suites.retain(|id| {
local_signature_schemes.iter().any(|algorithm| {
let signature_family_matches = match id {
CipherSuiteId::Tls_Ecdhe_Rsa_With_Aes_128_Gcm_Sha256
| CipherSuiteId::Tls_Ecdhe_Rsa_With_Aes_256_Cbc_Sha
| CipherSuiteId::Tls_Ecdhe_Rsa_With_ChaCha20_Poly1305_Sha256 => {
algorithm.signature
== crate::signature_hash_algorithm::SignatureAlgorithm::Rsa
}
_ => {
algorithm.signature
== crate::signature_hash_algorithm::SignatureAlgorithm::Ecdsa
}
};
signature_family_matches
&& algorithm
.crypto_scheme()
.is_ok_and(|scheme| signing_key.supports(scheme))
})
});
if local_cipher_suites.is_empty() {
return Err(Error::ErrNoAvailableCipherSuites);
}
}
let retransmit_interval = if self.flight_interval != Duration::from_secs(0) {
self.flight_interval
} else {
INITIAL_TICKER_INTERVAL
};
let maximum_transmission_unit = if self.mtu == 0 { DEFAULT_MTU } else { self.mtu };
let maximum_retransmit_number = DEFAULT_MAXIMUM_RETRANSMIT_NUMBER;
let replay_protection_window = if self.replay_protection_window == 0 {
DEFAULT_REPLAY_PROTECTION_WINDOW
} else {
self.replay_protection_window
};
let mut server_name = self.server_name.clone();
if is_client && server_name.is_empty() {
if let Some(remote_addr) = remote_addr {
server_name = remote_addr.ip().to_string();
} else {
warn!(
"conn.remote_addr is empty, please set explicitly server_name in Config! Use default \"localhost\" as server_name now"
);
"localhost".clone_into(&mut server_name);
}
}
let server_cert_verifier = if self.insecure_skip_verify {
None
} else {
let adapter = self.verifier_adapter.as_ref().ok_or_else(|| {
Error::Crypto("CA-chain verification requires a RustlsVerifierAdapter".to_owned())
})?;
let roots = if self.roots_cas.is_empty() {
gen_self_signed_root_cert()
} else {
self.roots_cas.clone()
};
Some(server_cert_verifier(Arc::new(roots), adapter)?)
};
let client_cert_verifier = if self.client_auth as u8
>= ClientAuthType::VerifyClientCertIfGiven as u8
{
let adapter = self.verifier_adapter.as_ref().ok_or_else(|| {
Error::Crypto(
"client-certificate verification requires a RustlsVerifierAdapter".to_owned(),
)
})?;
Some(
rustls::server::WebPkiClientVerifier::builder_with_provider(
Arc::new(self.client_cas.clone()),
adapter.provider.clone(),
)
.build()
.map_err(|err| Error::Other(format!("rustls client cert verifier: {err}")))?
as Arc<dyn ClientCertVerifier>,
)
} else {
None
};
Ok(HandshakeConfig {
crypto_provider,
local_psk_callback: self.psk.take(),
local_psk_identity_hint: self.psk_identity_hint.take(),
local_cipher_suites,
local_named_curves,
local_signature_schemes,
extended_master_secret: self.extended_master_secret,
local_srtp_protection_profiles: self.srtp_protection_profiles,
server_name,
client_auth: self.client_auth,
local_certificates: self.certificates,
name_to_certificate: Default::default(),
insecure_skip_verify: self.insecure_skip_verify,
insecure_verification: self.insecure_verification,
verify_peer_certificate: self.verify_peer_certificate.take(),
roots_cas: self.roots_cas,
server_cert_verifier,
client_cert_verifier,
retransmit_interval,
initial_epoch: 0,
maximum_transmission_unit,
maximum_retransmit_number,
replay_protection_window,
})
}
}
pub type VerifyPeerCertificateFn =
Arc<dyn (Fn(&[Vec<u8>], &[CertificateDer<'static>]) -> Result<()>) + Send + Sync>;
pub fn gen_self_signed_root_cert() -> rustls::RootCertStore {
#[cfg(any(feature = "crypto-ring", feature = "crypto-aws-lc-rs"))]
{
let mut certs = rustls::RootCertStore::empty();
certs
.add(
rcgen::generate_simple_self_signed(vec![])
.unwrap()
.cert
.der()
.to_owned(),
)
.unwrap();
certs
}
#[cfg(not(any(feature = "crypto-ring", feature = "crypto-aws-lc-rs")))]
{
rustls::RootCertStore::empty()
}
}
#[derive(Clone)]
pub struct HandshakeConfig {
pub(crate) crypto_provider: Arc<dyn RTCCryptoProvider>,
pub(crate) local_psk_callback: Option<PskCallback>,
pub(crate) local_psk_identity_hint: Option<Vec<u8>>,
pub(crate) local_cipher_suites: Vec<CipherSuiteId>, pub(crate) local_named_curves: Vec<NamedCurve>,
pub(crate) local_signature_schemes: Vec<SignatureHashAlgorithm>, pub(crate) extended_master_secret: ExtendedMasterSecretType, pub(crate) local_srtp_protection_profiles: Vec<SrtpProtectionProfile>, pub(crate) server_name: String,
pub(crate) client_auth: ClientAuthType, pub(crate) local_certificates: Vec<Certificate>,
pub(crate) name_to_certificate: HashMap<String, Certificate>,
pub(crate) insecure_skip_verify: bool,
pub(crate) insecure_verification: bool,
pub(crate) verify_peer_certificate: Option<VerifyPeerCertificateFn>,
pub(crate) roots_cas: rustls::RootCertStore,
pub(crate) server_cert_verifier: Option<Arc<dyn ServerCertVerifier>>,
pub(crate) client_cert_verifier: Option<Arc<dyn ClientCertVerifier>>,
pub(crate) retransmit_interval: Duration,
pub(crate) initial_epoch: u16,
pub(crate) maximum_transmission_unit: usize,
pub(crate) maximum_retransmit_number: usize,
pub(crate) replay_protection_window: usize,
}
impl fmt::Debug for HandshakeConfig {
fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt.debug_struct("HandshakeConfig<T>")
.field("crypto_provider", &self.crypto_provider.name())
.field("local_psk_identity_hint", &self.local_psk_identity_hint)
.field("local_cipher_suites", &self.local_cipher_suites)
.field("local_named_curves", &self.local_named_curves)
.field("local_signature_schemes", &self.local_signature_schemes)
.field("extended_master_secret", &self.extended_master_secret)
.field(
"local_srtp_protection_profiles",
&self.local_srtp_protection_profiles,
)
.field("server_name", &self.server_name)
.field("client_auth", &self.client_auth)
.field("local_certificates", &self.local_certificates)
.field("name_to_certificate", &self.name_to_certificate)
.field("insecure_skip_verify", &self.insecure_skip_verify)
.field("insecure_verification", &self.insecure_verification)
.field("roots_cas", &self.roots_cas)
.field("retransmit_interval", &self.retransmit_interval)
.field("initial_epoch", &self.initial_epoch)
.field("maximum_transmission_unit", &self.maximum_transmission_unit)
.field("maximum_retransmit_number", &self.maximum_retransmit_number)
.field("replay_protection_window", &self.replay_protection_window)
.finish()
}
}
impl HandshakeConfig {
pub(crate) fn new(crypto_provider: Arc<dyn RTCCryptoProvider>) -> Self {
Self {
crypto_provider,
local_psk_callback: None,
local_psk_identity_hint: None,
local_cipher_suites: vec![],
local_named_curves: vec![],
local_signature_schemes: vec![],
extended_master_secret: ExtendedMasterSecretType::Disable,
local_srtp_protection_profiles: vec![],
server_name: String::new(),
client_auth: ClientAuthType::NoClientCert,
local_certificates: vec![],
name_to_certificate: HashMap::new(),
insecure_skip_verify: false,
insecure_verification: false,
verify_peer_certificate: None,
roots_cas: rustls::RootCertStore::empty(),
server_cert_verifier: default_verifier_adapter().and_then(|adapter| {
server_cert_verifier(Arc::new(gen_self_signed_root_cert()), &adapter).ok()
}),
client_cert_verifier: None,
retransmit_interval: std::time::Duration::from_secs(0),
initial_epoch: 0,
maximum_transmission_unit: DEFAULT_MTU,
maximum_retransmit_number: DEFAULT_MAXIMUM_RETRANSMIT_NUMBER,
replay_protection_window: DEFAULT_REPLAY_PROTECTION_WINDOW,
}
}
pub(crate) fn provider(&self) -> &Arc<dyn RTCCryptoProvider> {
&self.crypto_provider
}
pub(crate) fn get_certificate(&self, server_name: &str) -> Result<Certificate> {
if self.local_certificates.is_empty() {
return Err(Error::ErrNoCertificates);
}
if self.local_certificates.len() == 1 {
return Ok(self.local_certificates[0].clone());
}
if server_name.is_empty() {
return Ok(self.local_certificates[0].clone());
}
let lower = server_name.to_lowercase();
let name = lower.trim_end_matches('.');
if let Some(cert) = self.name_to_certificate.get(name) {
return Ok(cert.clone());
}
let mut labels: Vec<&str> = name.split_terminator('.').collect();
for i in 0..labels.len() {
labels[i] = "*";
let candidate = labels.join(".");
if let Some(cert) = self.name_to_certificate.get(&candidate) {
return Ok(cert.clone());
}
}
Ok(self.local_certificates[0].clone())
}
}