use std::sync::Arc;
use async_trait::async_trait;
use dig_ip::{CandidateSource, DialConfig, LocalStack, PeerCandidates};
use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName};
use rustls::ClientConfig;
use tokio::net::TcpStream;
use tokio_rustls::TlsConnector;
use crate::config::LocalIdentity;
use crate::error::MethodError;
use crate::method::{MethodOutcome, TraversalKind};
use crate::mtls::{CapturedPeerId, PeerIdPinningVerifier};
use crate::peer::{PeerConnection, PeerTarget};
use crate::strategy::Dialer;
#[derive(Debug, Clone, Copy)]
pub struct HappyEyeballsConfig {
pub per_attempt_timeout: std::time::Duration,
pub stagger: std::time::Duration,
}
impl Default for HappyEyeballsConfig {
fn default() -> Self {
HappyEyeballsConfig {
per_attempt_timeout: std::time::Duration::from_secs(10),
stagger: std::time::Duration::from_millis(250),
}
}
}
impl From<HappyEyeballsConfig> for DialConfig {
fn from(cfg: HappyEyeballsConfig) -> DialConfig {
DialConfig {
per_attempt_timeout: cfg.per_attempt_timeout,
attempt_delay: cfg.stagger,
}
}
}
pub fn candidates_from_outcome(outcome: &MethodOutcome) -> PeerCandidates {
let source = match outcome.kind {
TraversalKind::HolePunch | TraversalKind::Relayed => CandidateSource::RelayIntroduction,
TraversalKind::Direct
| TraversalKind::Upnp
| TraversalKind::NatPmp
| TraversalKind::Pcp => CandidateSource::ListenAddr,
};
let mut candidates = PeerCandidates::new();
candidates.extend(outcome.dial_addrs.iter().copied(), source);
candidates
}
#[derive(Debug, Clone)]
pub struct MtlsDialer {
identity: LocalIdentity,
happy_eyeballs: HappyEyeballsConfig,
local_stack: Option<LocalStack>,
}
impl MtlsDialer {
pub fn new(identity: LocalIdentity) -> Self {
MtlsDialer {
identity,
happy_eyeballs: HappyEyeballsConfig::default(),
local_stack: None,
}
}
pub fn with_happy_eyeballs(mut self, config: HappyEyeballsConfig) -> Self {
self.happy_eyeballs = config;
self
}
pub fn with_local_stack(mut self, stack: LocalStack) -> Self {
self.local_stack = Some(stack);
self
}
fn client_config(
&self,
expected: crate::identity::PeerId,
captured: CapturedPeerId,
) -> Result<ClientConfig, String> {
let cert = CertificateDer::from(self.identity.cert_der.clone());
let key = PrivateKeyDer::try_from(self.identity.key_der.to_vec())
.map_err(|e| format!("invalid private key: {e}"))?;
let verifier = Arc::new(PeerIdPinningVerifier::new(Some(expected), captured));
ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(verifier)
.with_client_auth_cert(vec![cert], key)
.map_err(|e| format!("client cert config: {e}"))
}
}
#[async_trait]
impl Dialer for MtlsDialer {
async fn dial(
&self,
peer: &PeerTarget,
outcome: &MethodOutcome,
) -> Result<PeerConnection, MethodError> {
let kind = outcome.kind;
let local = self.local_stack.unwrap_or_else(LocalStack::cached);
let candidates = candidates_from_outcome(outcome);
let winner = dig_ip::connect(
&local,
&candidates,
self.happy_eyeballs.into(),
|addr| async move {
TcpStream::connect(addr)
.await
.map_err(|e| format!("tcp connect {addr}: {e}"))
},
)
.await
.map_err(|e| MethodError::failed(kind, e.to_string()))?;
let tcp = winner.conn;
let addr = winner.addr;
let captured = CapturedPeerId::default();
let config = self
.client_config(peer.peer_id, captured.clone())
.map_err(|e| MethodError::failed(kind, e))?;
let connector = TlsConnector::from(Arc::new(config));
let server_name = ServerName::try_from("peer.dig.invalid")
.map_err(|e| MethodError::failed(kind, format!("server name: {e}")))?;
let tls = connector
.connect(server_name, tcp)
.await
.map_err(|e| classify_tls_error(kind, &e))?;
let verified = captured
.get()
.ok_or_else(|| MethodError::failed(kind, "peer presented no certificate"))?;
let session = crate::mux::PeerSession::client(tls);
Ok(PeerConnection {
peer_id: verified,
method: kind,
remote_addr: addr,
session,
})
}
}
fn classify_tls_error(kind: TraversalKind, e: &std::io::Error) -> MethodError {
let msg = e.to_string();
MethodError::failed(kind, format!("mtls handshake: {msg}"))
}