use std::net::SocketAddr;
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::io::{AsyncRead, AsyncWrite};
use tokio::net::TcpStream;
use tokio_rustls::TlsConnector;
use crate::error::MethodError;
use crate::method::relayed::RelayedDialer;
use crate::method::{MethodOutcome, TraversalKind};
use crate::peer::{PeerConnection, PeerTarget};
use crate::strategy::Dialer;
use crate::tunnel::RelayTunnelStream;
#[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(Clone)]
pub struct MtlsDialer {
node: Arc<NodeCert>,
happy_eyeballs: HappyEyeballsConfig,
binding_policy: BindingPolicy,
local_stack: Option<LocalStack>,
relayed: Option<Arc<dyn RelayedDialer>>,
}
impl MtlsDialer {
pub fn new(node: Arc<NodeCert>) -> Self {
MtlsDialer {
node,
happy_eyeballs: HappyEyeballsConfig::default(),
binding_policy: BindingPolicy::default(),
local_stack: None,
relayed: 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
}
pub fn with_relayed_dialer(mut self, relayed: Arc<dyn RelayedDialer>) -> Self {
self.relayed = Some(relayed);
self
}
async fn handshake_over<S>(
&self,
peer: &PeerTarget,
kind: TraversalKind,
stream: S,
remote_addr: SocketAddr,
) -> Result<PeerConnection, MethodError>
where
S: AsyncRead + AsyncWrite + Send + Unpin + 'static,
{
let client_tls =
dig_tls::client_config_spki_pinned(&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, stream)
.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,
peer_bls_pub: captured_bls.get(),
session,
})
}
async fn dial_relayed(
&self,
peer: &PeerTarget,
outcome: &MethodOutcome,
) -> Result<PeerConnection, MethodError> {
let kind = TraversalKind::Relayed;
let relayed = self.relayed.as_ref().ok_or_else(|| {
MethodError::failed(kind, "no relay data-plane wired for the relayed tier")
})?;
let tunnel = relayed
.open_dial_tunnel(&peer.peer_id.to_hex(), &peer.network_id)
.await
.map_err(|e| MethodError::failed(kind, e))?;
let remote_addr = outcome.dial_addr().unwrap_or(relayed.relay_endpoint());
let stream = RelayTunnelStream::new(tunnel);
self.handshake_over(peer, kind, stream, remote_addr).await
}
}
impl std::fmt::Debug for MtlsDialer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtlsDialer")
.field("happy_eyeballs", &self.happy_eyeballs)
.field("binding_policy", &self.binding_policy)
.field("local_stack", &self.local_stack)
.field(
"relayed",
&self.relayed.as_ref().map(|_| "<relay data-plane>"),
)
.finish_non_exhaustive()
}
}
#[async_trait]
impl Dialer for MtlsDialer {
async fn dial(
&self,
peer: &PeerTarget,
outcome: &MethodOutcome,
) -> Result<PeerConnection, MethodError> {
let kind = outcome.kind;
if kind == TraversalKind::Relayed {
return self.dial_relayed(peer, outcome).await;
}
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()))?;
self.handshake_over(peer, kind, winner.conn, winner.addr)
.await
}
}
fn classify_tls_error(kind: TraversalKind, e: &std::io::Error) -> MethodError {
let msg = e.to_string();
MethodError::failed(kind, format!("mtls handshake: {msg}"))
}