use std::sync::Arc;
use async_trait::async_trait;
use dig_ip::{CandidateSource, DialConfig, LocalStack, PeerCandidates};
use dig_tls::{BindingPolicy, NodeCert};
use rustls_pki_types::ServerName;
use tokio::net::TcpStream;
use tokio_rustls::TlsConnector;
use crate::error::MethodError;
use crate::method::{MethodOutcome, TraversalKind};
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 {
node: Arc<NodeCert>,
happy_eyeballs: HappyEyeballsConfig,
binding_policy: BindingPolicy,
local_stack: Option<LocalStack>,
}
impl MtlsDialer {
pub fn new(node: Arc<NodeCert>) -> Self {
MtlsDialer {
node,
happy_eyeballs: HappyEyeballsConfig::default(),
binding_policy: BindingPolicy::default(),
local_stack: None,
}
}
pub fn with_happy_eyeballs(mut self, config: HappyEyeballsConfig) -> Self {
self.happy_eyeballs = config;
self
}
pub fn with_binding_policy(mut self, policy: BindingPolicy) -> Self {
self.binding_policy = policy;
self
}
pub fn with_local_stack(mut self, stack: LocalStack) -> Self {
self.local_stack = Some(stack);
self
}
}
#[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 client_tls =
dig_tls::client_config(&self.node, Some(peer.peer_id), self.binding_policy)
.map_err(|e| MethodError::failed(kind, format!("client cert config: {e}")))?;
let captured = client_tls.captured_peer_id;
let captured_bls = client_tls.captured_bls;
let connector = TlsConnector::from(client_tls.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,
peer_bls_pub: captured_bls.get(),
session,
})
}
}
fn classify_tls_error(kind: TraversalKind, e: &std::io::Error) -> MethodError {
let msg = e.to_string();
MethodError::failed(kind, format!("mtls handshake: {msg}"))
}