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 peer.peer_id == self.node.peer_id() {
return Err(MethodError::failed(
kind,
"refusing self-dial (target peer_id == local peer_id)",
));
}
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}"))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::method::MethodOutcome;
use crate::peer::PeerTarget;
use std::net::SocketAddr;
fn node_cert() -> Arc<NodeCert> {
use sha2::{Digest, Sha256};
let seed: [u8; 32] = Sha256::digest(b"dig-nat/dialer/self-dial-test").into();
let bls_sk = dig_tls::bls::SecretKey::from_seed(&seed);
Arc::new(NodeCert::generate_signed(&bls_sk).unwrap())
}
fn addr(s: &str) -> SocketAddr {
s.parse().unwrap()
}
#[tokio::test]
async fn direct_dial_to_own_peer_id_is_refused_before_racing_candidates() {
let node = node_cert();
let dialer = MtlsDialer::new(node.clone());
let self_peer = PeerTarget::with_addrs(
node.peer_id(),
vec![addr("127.0.0.1:1"), addr("[::1]:1")],
"DIG_MAINNET",
);
let outcome = MethodOutcome::candidates(
TraversalKind::Direct,
vec![addr("127.0.0.1:1"), addr("[::1]:1")],
);
let err = dialer.dial(&self_peer, &outcome).await.unwrap_err();
assert!(
err.reason.contains("self-dial"),
"a Direct dial to our own peer_id must be refused as a self-dial, got: {}",
err.reason
);
let outcome_rev = MethodOutcome::candidates(
TraversalKind::Direct,
vec![addr("[::1]:1"), addr("127.0.0.1:1")],
);
let err_rev = dialer.dial(&self_peer, &outcome_rev).await.unwrap_err();
assert!(err_rev.reason.contains("self-dial"), "{}", err_rev.reason);
}
#[tokio::test]
async fn relayed_dial_to_own_peer_id_is_refused_at_the_chokepoint() {
let node = node_cert();
let dialer = MtlsDialer::new(node.clone());
let self_peer = PeerTarget::relay_only(node.peer_id(), "DIG_MAINNET");
let outcome = MethodOutcome::candidates(TraversalKind::Relayed, vec![]);
let err = dialer.dial(&self_peer, &outcome).await.unwrap_err();
assert!(err.reason.contains("self-dial"), "{}", err.reason);
}
#[tokio::test]
async fn dial_to_a_different_peer_id_is_not_refused_as_self() {
let node = node_cert();
let dialer = MtlsDialer::new(node.clone());
let holder = PeerTarget::with_addr(
dig_tls::PeerId::from_bytes([0x5A; 32]),
addr("127.0.0.1:1"),
"DIG_MAINNET",
);
let outcome = MethodOutcome::candidates(TraversalKind::Direct, vec![addr("127.0.0.1:1")]);
let err = dialer.dial(&holder, &outcome).await.unwrap_err();
assert!(
!err.reason.contains("self-dial"),
"a dial to a different peer must not be refused as a self-dial, got: {}",
err.reason
);
}
}