use std::future::Future;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use arc_swap::ArcSwap;
use futures::future::{select, Either};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, ReadBuf};
use tokio::sync::{watch, Mutex, Notify};
use crate::dialer::MtlsDialer;
use crate::error::NatError;
use crate::method::{MethodOutcome, TraversalKind};
use crate::mux::{
AvailabilityRequest, AvailabilityResponse, ClosedHandle, PeerSession, PeerStream, RangeRequest,
};
use crate::peer::{PeerConnection, PeerTarget};
use crate::strategy::{self, Dialer};
use crate::{NatConfig, NatRuntime, NodeCert, PeerId};
type DialFuture = Pin<Box<dyn Future<Output = Result<PeerConnection, NatError>> + Send>>;
type Establisher = Arc<dyn Fn() -> DialFuture + Send + Sync>;
struct TransportSlot {
session: Mutex<PeerSession>,
method: TraversalKind,
remote_addr: SocketAddr,
peer_bls_pub: Option<[u8; 48]>,
closed: ClosedHandle,
outstanding: AtomicUsize,
drained: Notify,
}
impl TransportSlot {
fn from_conn(conn: PeerConnection) -> Arc<TransportSlot> {
let closed = conn.session.closed_handle();
Arc::new(TransportSlot {
session: Mutex::new(conn.session),
method: conn.method,
remote_addr: conn.remote_addr,
peer_bls_pub: conn.peer_bls_pub,
closed,
outstanding: AtomicUsize::new(0),
drained: Notify::new(),
})
}
}
pub struct FastPeerConnection {
peer_id: PeerId,
active: Arc<ArcSwap<TransportSlot>>,
events: watch::Sender<TraversalKind>,
_guard: PromotionGuard,
}
pub struct FastPeerStream {
inner: PeerStream,
slot: Arc<TransportSlot>,
}
impl Drop for FastPeerStream {
fn drop(&mut self) {
if self.slot.outstanding.fetch_sub(1, Ordering::AcqRel) == 1 {
self.slot.drained.notify_waiters();
}
}
}
impl AsyncRead for FastPeerStream {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl AsyncWrite for FastPeerStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
}
impl FastPeerConnection {
pub fn peer_id(&self) -> PeerId {
self.peer_id
}
pub fn current_method(&self) -> TraversalKind {
self.active.load().method
}
pub fn remote_addr(&self) -> SocketAddr {
self.active.load().remote_addr
}
pub fn subscribe(&self) -> watch::Receiver<TraversalKind> {
self.events.subscribe()
}
pub async fn open_stream(&self) -> std::io::Result<FastPeerStream> {
let slot = self.active.load_full();
let stream = {
let mut session = slot.session.lock().await;
session.open_stream().await?
};
slot.outstanding.fetch_add(1, Ordering::AcqRel);
Ok(FastPeerStream {
inner: stream,
slot,
})
}
pub async fn open_range_stream(&self, req: &RangeRequest) -> std::io::Result<FastPeerStream> {
let mut stream = self.open_stream().await?;
stream.write_all(&req.encode()).await?;
stream.flush().await?;
Ok(stream)
}
pub async fn query_availability(
&self,
items: Vec<crate::mux::AvailabilityItem>,
) -> std::io::Result<AvailabilityResponse> {
let mut stream = self.open_stream().await?;
stream
.write_all(&AvailabilityRequest { items }.encode())
.await?;
stream.flush().await?;
AvailabilityResponse::decode(&mut stream).await
}
}
impl std::fmt::Debug for FastPeerConnection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FastPeerConnection")
.field("peer_id", &self.peer_id)
.field("method", &self.current_method())
.field("remote_addr", &self.remote_addr())
.finish_non_exhaustive()
}
}
struct PromotionGuard {
handle: tokio::task::JoinHandle<()>,
}
impl Drop for PromotionGuard {
fn drop(&mut self) {
self.handle.abort();
}
}
pub async fn connect_fast(
peer: &PeerTarget,
node: &Arc<NodeCert>,
config: &NatConfig,
runtime: &NatRuntime,
) -> Result<FastPeerConnection, NatError> {
let mut direct_config = config.clone();
direct_config
.enabled_methods
.retain(|k| *k != TraversalKind::Relayed);
let direct_methods = crate::compose_ladder(&direct_config, runtime);
let direct_dialer =
Arc::new(MtlsDialer::new(Arc::clone(node)).with_binding_policy(config.binding_policy));
let direct: Establisher = {
let peer = peer.clone();
let timeout = config.per_method_timeout;
Arc::new(move || {
let peer = peer.clone();
let dialer = Arc::clone(&direct_dialer);
let methods = direct_methods.clone();
Box::pin(async move {
strategy::connect_with_strategy(&peer, methods, dialer.as_ref(), timeout).await
})
})
};
let relayed: Option<Establisher> = runtime.relayed.as_ref().map(|relayed_dialer| {
let relayed_dialer = Arc::clone(relayed_dialer);
let node = Arc::clone(node);
let peer = peer.clone();
let binding = config.binding_policy;
let endpoint = relayed_dialer.relay_endpoint();
let est: Establisher = Arc::new(move || {
let dialer = MtlsDialer::new(Arc::clone(&node))
.with_binding_policy(binding)
.with_relayed_dialer(Arc::clone(&relayed_dialer));
let peer = peer.clone();
Box::pin(async move {
let outcome = MethodOutcome::single(TraversalKind::Relayed, endpoint);
dialer
.dial(&peer, &outcome)
.await
.map_err(|e| NatError::AllMethodsFailed(vec![e]))
})
});
est
});
connect_fast_with(
peer.peer_id,
direct,
relayed,
config.fast_connect_grace,
config.per_method_timeout,
)
.await
}
async fn connect_fast_with(
expected_peer_id: PeerId,
direct: Establisher,
relayed: Option<Establisher>,
grace_cap: Duration,
probe_timeout: Duration,
) -> Result<FastPeerConnection, NatError> {
let direct_fut = direct();
let Some(relayed) = relayed else {
let conn = direct_fut.await?;
return Ok(build(
expected_peer_id,
conn,
GuardPlan::none(),
grace_cap,
probe_timeout,
));
};
let relayed_fut = relayed();
match select(relayed_fut, direct_fut).await {
Either::Left((relayed_res, direct_fut)) => match relayed_res {
Ok(relayed_conn) => {
let plan = GuardPlan {
promote_from: Some(direct_fut),
fallback: Some(relayed),
};
Ok(build(
expected_peer_id,
relayed_conn,
plan,
grace_cap,
probe_timeout,
))
}
Err(relayed_err) => match direct_fut.await {
Ok(direct_conn) => Ok(build(
expected_peer_id,
direct_conn,
GuardPlan::fallback_only(Some(relayed)),
grace_cap,
probe_timeout,
)),
Err(direct_err) => Err(merge_errors(relayed_err, direct_err)),
},
},
Either::Right((direct_res, relayed_fut)) => match direct_res {
Ok(direct_conn) => {
drop(relayed_fut);
Ok(build(
expected_peer_id,
direct_conn,
GuardPlan::fallback_only(Some(relayed)),
grace_cap,
probe_timeout,
))
}
Err(direct_err) => match relayed_fut.await {
Ok(relayed_conn) => {
let plan = GuardPlan {
promote_from: Some(direct()),
fallback: Some(relayed),
};
Ok(build(
expected_peer_id,
relayed_conn,
plan,
grace_cap,
probe_timeout,
))
}
Err(relayed_err) => Err(merge_errors(relayed_err, direct_err)),
},
},
}
}
struct GuardPlan {
promote_from: Option<DialFuture>,
fallback: Option<Establisher>,
}
impl GuardPlan {
fn none() -> Self {
GuardPlan {
promote_from: None,
fallback: None,
}
}
fn fallback_only(fallback: Option<Establisher>) -> Self {
GuardPlan {
promote_from: None,
fallback,
}
}
}
fn build(
peer_id: PeerId,
initial: PeerConnection,
plan: GuardPlan,
grace_cap: Duration,
probe_timeout: Duration,
) -> FastPeerConnection {
let slot = TransportSlot::from_conn(initial);
let (events, _rx) = watch::channel(slot.method);
let active = Arc::new(ArcSwap::from(slot));
let handle = tokio::spawn(run_guard(
peer_id,
Arc::clone(&active),
events.clone(),
plan,
grace_cap,
probe_timeout,
));
FastPeerConnection {
peer_id,
active,
events,
_guard: PromotionGuard { handle },
}
}
async fn run_guard(
peer_id: PeerId,
active: Arc<ArcSwap<TransportSlot>>,
events: watch::Sender<TraversalKind>,
plan: GuardPlan,
grace_cap: Duration,
probe_timeout: Duration,
) {
if let Some(promote_from) = plan.promote_from {
if let Ok(direct_conn) = promote_from.await {
try_promote(
peer_id,
&active,
&events,
direct_conn,
grace_cap,
probe_timeout,
)
.await;
}
}
let Some(fallback) = plan.fallback else {
return;
};
let mut rapid_deaths: u32 = 0;
loop {
let closed = active.load().closed.clone();
let established_at = tokio::time::Instant::now();
closed.closed().await;
if established_at.elapsed() >= FALLBACK_STABILITY {
rapid_deaths = 0;
} else {
rapid_deaths = rapid_deaths.saturating_add(1);
}
let backoff = fallback_backoff(rapid_deaths, FALLBACK_BACKOFF_BASE, FALLBACK_BACKOFF_CAP);
if !backoff.is_zero() {
tokio::time::sleep(backoff).await;
}
match fallback().await {
Ok(conn) => {
let slot = TransportSlot::from_conn(conn);
let method = slot.method;
active.store(slot);
let _ = events.send(method);
}
Err(_) => return,
}
}
}
const FALLBACK_BACKOFF_BASE: Duration = Duration::from_millis(50);
const FALLBACK_BACKOFF_CAP: Duration = Duration::from_secs(5);
const FALLBACK_STABILITY: Duration = Duration::from_secs(10);
fn fallback_backoff(rapid_deaths: u32, base: Duration, cap: Duration) -> Duration {
if rapid_deaths == 0 {
return Duration::ZERO;
}
let base_ms = base.as_millis() as u64;
let shifted = base_ms.checked_shl(rapid_deaths - 1).unwrap_or(u64::MAX);
Duration::from_millis(shifted).clamp(base, cap)
}
async fn try_promote(
peer_id: PeerId,
active: &Arc<ArcSwap<TransportSlot>>,
events: &watch::Sender<TraversalKind>,
mut direct_conn: PeerConnection,
grace_cap: Duration,
probe_timeout: Duration,
) {
let relayed_slot = active.load_full();
if direct_conn.peer_id != peer_id || direct_conn.peer_bls_pub != relayed_slot.peer_bls_pub {
tracing::warn!(
"fast-connect: direct path identity mismatch — promotion refused, staying relayed"
);
return;
}
match tokio::time::timeout(probe_timeout, direct_conn.query_availability(vec![])).await {
Ok(Ok(_)) => {}
Ok(Err(_)) => {
tracing::warn!(
"fast-connect: direct path failed the availability probe — staying relayed"
);
return;
}
Err(_) => {
tracing::warn!(
"fast-connect: direct path availability probe timed out — staying relayed"
);
return;
}
}
let direct_slot = TransportSlot::from_conn(direct_conn);
let method = direct_slot.method;
active.store(direct_slot);
let _ = events.send(method);
tracing::info!(?method, "fast-connect: promoted to a direct transport");
tokio::spawn(drain_then_drop(relayed_slot, grace_cap));
}
async fn drain_then_drop(slot: Arc<TransportSlot>, grace_cap: Duration) {
let deadline = tokio::time::sleep(grace_cap);
tokio::pin!(deadline);
loop {
if slot.outstanding.load(Ordering::Acquire) == 0 {
break;
}
let drained = slot.drained.notified();
if slot.outstanding.load(Ordering::Acquire) == 0 {
break;
}
tokio::select! {
_ = drained => {}
_ = &mut deadline => break,
}
}
drop(slot);
}
fn merge_errors(relayed: NatError, direct: NatError) -> NatError {
let mut failures = Vec::new();
for e in [relayed, direct] {
if let NatError::AllMethodsFailed(mut fs) = e {
failures.append(&mut fs);
}
}
NatError::AllMethodsFailed(failures)
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
use sha2::{Digest, Sha256};
use tokio::io::AsyncWriteExt;
use tokio_rustls::TlsAcceptor;
use crate::method::relayed::ReservationRelayedTransport;
use crate::mux::{
AvailabilityAnswer, AvailabilityItem, AvailabilityRequest, AvailabilityResponse,
};
use crate::relay::{loopback_reservation_pair, RelayStatus};
use crate::tunnel::RelayTunnelStream;
use crate::{BindingPolicy, MethodError};
use dig_tls::bls::SecretKey;
const NET: &str = "DIG_MAINNET";
const RELAY_ENDPOINT: &str = "127.0.0.1:3478";
fn test_bls_sk(label: &str) -> SecretKey {
let seed: [u8; 32] = Sha256::digest(label.as_bytes()).into();
SecretKey::from_seed(&seed)
}
fn test_node(label: &str) -> Arc<NodeCert> {
Arc::new(NodeCert::generate_signed(&test_bls_sk(label)).expect("generate node cert"))
}
fn serve_availability<S>(acceptor: TlsAcceptor, stream: S, tag: u64, kill: Option<Arc<Notify>>)
where
S: AsyncRead + AsyncWrite + Send + Unpin + 'static,
{
tokio::spawn(async move {
let Ok(tls) = acceptor.accept(stream).await else {
return;
};
let mut session = PeerSession::server(tls);
loop {
let accepted = match &kill {
Some(kill) => tokio::select! {
s = session.accept_stream() => s,
_ = kill.notified() => return, },
None => session.accept_stream().await,
};
let Some(mut s) = accepted else { return };
tokio::spawn(async move {
if let Ok(req) = AvailabilityRequest::decode(&mut s).await {
let resp = AvailabilityResponse {
items: req
.items
.iter()
.map(|_| AvailabilityAnswer {
available: true,
roots: None,
total_length: Some(tag),
chunk_count: Some(1),
complete: Some(true),
})
.collect(),
};
let _ = s.write_all(&resp.encode()).await;
let _ = s.shutdown().await;
}
});
}
});
}
fn serve_blackhole<S>(acceptor: TlsAcceptor, stream: S)
where
S: AsyncRead + AsyncWrite + Send + Unpin + 'static,
{
tokio::spawn(async move {
let Ok(tls) = acceptor.accept(stream).await else {
return;
};
let _session = PeerSession::server(tls);
std::future::pending::<()>().await;
});
}
fn blackhole_direct_establisher(
client: &Arc<NodeCert>,
server: &Arc<NodeCert>,
delay: Duration,
) -> Establisher {
let node = Arc::clone(client);
let server = Arc::clone(server);
let server_id = server.peer_id();
Arc::new(move || {
let node = Arc::clone(&node);
let server = Arc::clone(&server);
Box::pin(async move {
if !delay.is_zero() {
tokio::time::sleep(delay).await;
}
let (client_io, server_io) = tokio::io::duplex(64 * 1024);
let server_tls = dig_tls::server_config(&server, BindingPolicy::Opportunistic)
.expect("server config")
.config;
serve_blackhole(TlsAcceptor::from(server_tls), server_io);
let client_cfg =
dig_tls::client_config(&node, Some(server_id), BindingPolicy::Opportunistic)
.expect("client config");
let captured = client_cfg.captured_peer_id;
let captured_bls = client_cfg.captured_bls;
let connector = tokio_rustls::TlsConnector::from(client_cfg.config);
let sni = rustls_pki_types::ServerName::try_from("peer.dig.invalid").unwrap();
let tls = connector.connect(sni, client_io).await.map_err(|e| {
NatError::AllMethodsFailed(vec![MethodError::failed(
TraversalKind::Direct,
format!("mtls handshake: {e}"),
)])
})?;
let verified = captured.get().expect("peer presented a cert");
Ok(PeerConnection {
peer_id: verified,
method: TraversalKind::Direct,
remote_addr: "203.0.113.9:4444".parse().unwrap(),
peer_bls_pub: captured_bls.get(),
session: PeerSession::client(tls),
})
})
})
}
fn mismatched_bls_establisher(
client: &Arc<NodeCert>,
server: &Arc<NodeCert>,
tag: u64,
delay: Duration,
) -> Establisher {
let inner = direct_establisher(client, server, tag, delay);
Arc::new(move || {
let inner = Arc::clone(&inner);
Box::pin(async move {
let mut conn = inner().await?;
conn.peer_bls_pub = Some([0xAB; 48]);
Ok(conn)
})
})
}
fn relayed_establisher(
client: &Arc<NodeCert>,
server: &Arc<NodeCert>,
tag: u64,
) -> (Establisher, Arc<RelayStatus>) {
let client_hex = client.peer_id().to_hex();
let server_hex = server.peer_id().to_hex();
let (client_status, server_status) = loopback_reservation_pair(&client_hex, &server_hex);
let server_tunnel = server_status
.open_tunnel(&client_hex, NET)
.expect("server opens relay tunnel");
let server_tls = dig_tls::server_config(server, BindingPolicy::Opportunistic)
.expect("server config")
.config;
serve_availability(
TlsAcceptor::from(server_tls),
RelayTunnelStream::new(server_tunnel),
tag,
None,
);
let server_id = server.peer_id();
let node = Arc::clone(client);
let transport = Arc::new(ReservationRelayedTransport::new(
Arc::clone(&client_status),
RELAY_ENDPOINT.parse().unwrap(),
));
let est: Establisher = Arc::new(move || {
let dialer = MtlsDialer::new(Arc::clone(&node))
.with_binding_policy(BindingPolicy::Opportunistic)
.with_relayed_dialer(Arc::clone(&transport) as Arc<_>);
Box::pin(async move {
let peer = PeerTarget::relay_only(server_id, NET);
let outcome =
MethodOutcome::single(TraversalKind::Relayed, RELAY_ENDPOINT.parse().unwrap());
dialer
.dial(&peer, &outcome)
.await
.map_err(|e| NatError::AllMethodsFailed(vec![e]))
})
});
(est, client_status)
}
fn duplex_establisher(
client: &Arc<NodeCert>,
server: &Arc<NodeCert>,
method: TraversalKind,
tag: u64,
delay: Duration,
kill: Option<Arc<Notify>>,
) -> Establisher {
let node = Arc::clone(client);
let server = Arc::clone(server);
let server_id = server.peer_id();
Arc::new(move || {
let node = Arc::clone(&node);
let server = Arc::clone(&server);
let kill = kill.clone();
Box::pin(async move {
if !delay.is_zero() {
tokio::time::sleep(delay).await;
}
let (client_io, server_io) = tokio::io::duplex(64 * 1024);
let server_tls = dig_tls::server_config(&server, BindingPolicy::Opportunistic)
.expect("server config")
.config;
serve_availability(TlsAcceptor::from(server_tls), server_io, tag, kill);
let client_cfg =
dig_tls::client_config(&node, Some(server_id), BindingPolicy::Opportunistic)
.expect("client config");
let captured = client_cfg.captured_peer_id;
let captured_bls = client_cfg.captured_bls;
let connector = tokio_rustls::TlsConnector::from(client_cfg.config);
let sni = rustls_pki_types::ServerName::try_from("peer.dig.invalid").unwrap();
let tls = connector.connect(sni, client_io).await.map_err(|e| {
NatError::AllMethodsFailed(vec![MethodError::failed(
method,
format!("mtls handshake: {e}"),
)])
})?;
let verified = captured.get().expect("peer presented a cert");
Ok(PeerConnection {
peer_id: verified,
method,
remote_addr: "203.0.113.9:4444".parse().unwrap(),
peer_bls_pub: captured_bls.get(),
session: PeerSession::client(tls),
})
})
})
}
fn direct_establisher(
client: &Arc<NodeCert>,
server: &Arc<NodeCert>,
tag: u64,
delay: Duration,
) -> Establisher {
duplex_establisher(client, server, TraversalKind::Direct, tag, delay, None)
}
#[tokio::test]
async fn returns_first_usable_relayed_before_slow_direct() {
let client = test_node("fc/1/client");
let server = test_node("fc/1/server");
let (relayed, _status) = relayed_establisher(&client, &server, 11);
let direct = direct_establisher(&client, &server, 22, Duration::from_secs(30));
let conn = tokio::time::timeout(
Duration::from_secs(5),
connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_millis(200),
Duration::from_secs(5),
),
)
.await
.expect("connect_fast returns before the slow direct completes")
.expect("relayed lands first");
assert_eq!(conn.current_method(), TraversalKind::Relayed);
assert_eq!(conn.peer_id(), server.peer_id());
assert_eq!(conn.remote_addr(), RELAY_ENDPOINT.parse().unwrap());
assert!(format!("{conn:?}").contains("FastPeerConnection"));
let resp = conn.query_availability(vec![avail_item()]).await.unwrap();
assert_eq!(resp.items[0].total_length, Some(11));
let _range = conn
.open_range_stream(&RangeRequest::resource(
"aa".repeat(32),
"cc".repeat(32),
0,
8,
))
.await
.unwrap();
}
#[tokio::test]
async fn connect_fast_direct_only_with_no_methods_errors() {
let client = test_node("fc/pub/client");
let peer = PeerTarget::relay_only(test_node("fc/pub/server").peer_id(), NET);
let config = NatConfig::builder().enabled_methods(vec![]).build();
let runtime = NatRuntime::default(); let err = connect_fast(&peer, &client, &config, &runtime)
.await
.unwrap_err();
assert!(matches!(err, NatError::NoMethodsEnabled));
}
#[tokio::test]
async fn promotes_to_direct_without_losing_inflight_relayed_stream() {
let client = test_node("fc/2/client");
let server = test_node("fc/2/server");
let (relayed, _status) = relayed_establisher(&client, &server, 11);
let direct = direct_establisher(&client, &server, 22, Duration::from_millis(150));
let conn = connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_secs(5),
Duration::from_secs(5),
)
.await
.expect("relayed lands first");
assert_eq!(conn.current_method(), TraversalKind::Relayed);
let mut pre = conn.open_stream().await.unwrap();
pre.write_all(
&AvailabilityRequest {
items: vec![avail_item()],
}
.encode(),
)
.await
.unwrap();
pre.flush().await.unwrap();
let mut rx = conn.subscribe();
tokio::time::timeout(Duration::from_secs(5), async {
while *rx.borrow_and_update() != TraversalKind::Direct {
rx.changed().await.unwrap();
}
})
.await
.expect("promoted to Direct");
assert_eq!(conn.current_method(), TraversalKind::Direct);
let pre_resp = AvailabilityResponse::decode(&mut pre).await.unwrap();
assert_eq!(
pre_resp.items[0].total_length,
Some(11),
"in-flight stream stayed relayed"
);
let post = conn.query_availability(vec![avail_item()]).await.unwrap();
assert_eq!(post.items[0].total_length, Some(22), "new stream is direct");
}
#[tokio::test]
async fn drains_and_releases_relay_tunnel_after_promotion() {
let client = test_node("fc/3/client");
let server = test_node("fc/3/server");
let (relayed, status) = relayed_establisher(&client, &server, 11);
let direct = direct_establisher(&client, &server, 22, Duration::from_millis(100));
let server_hex = server.peer_id().to_hex();
let conn = connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_millis(100),
Duration::from_secs(5),
)
.await
.unwrap();
let mut rx = conn.subscribe();
tokio::time::timeout(Duration::from_secs(5), async {
while *rx.borrow_and_update() != TraversalKind::Direct {
rx.changed().await.unwrap();
}
})
.await
.expect("promoted");
wait_until(Duration::from_secs(3), || {
!status.open_tunnel_exists(&server_hex)
})
.await;
assert!(
!status.open_tunnel_exists(&server_hex),
"per-peer relay tunnel released after promotion+drain"
);
assert!(
status.is_connected(),
"the relay reservation stays Connected"
);
}
#[tokio::test]
async fn stays_relayed_when_direct_never_lands() {
let client = test_node("fc/4/client");
let server = test_node("fc/4/server");
let (relayed, status) = relayed_establisher(&client, &server, 11);
let direct: Establisher = Arc::new(|| {
Box::pin(async {
Err(NatError::AllMethodsFailed(vec![MethodError::failed(
TraversalKind::Direct,
"no direct path",
)]))
})
});
let conn = connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_millis(200),
Duration::from_secs(5),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(200)).await;
assert_eq!(conn.current_method(), TraversalKind::Relayed);
let resp = conn.query_availability(vec![avail_item()]).await.unwrap();
assert_eq!(resp.items[0].total_length, Some(11));
assert!(status.is_connected());
}
#[tokio::test]
async fn refuses_promotion_on_identity_mismatch() {
let client = test_node("fc/6/client");
let server = test_node("fc/6/server");
let impostor = test_node("fc/6/impostor");
let (relayed, _status) = relayed_establisher(&client, &server, 11);
let direct = direct_establisher(&client, &impostor, 22, Duration::from_millis(100));
let conn = connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_millis(200),
Duration::from_secs(5),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(400)).await;
assert_eq!(
conn.current_method(),
TraversalKind::Relayed,
"promotion refused on identity mismatch"
);
let resp = conn.query_availability(vec![avail_item()]).await.unwrap();
assert_eq!(resp.items[0].total_length, Some(11), "still relayed");
}
#[tokio::test]
async fn falls_back_to_relayed_when_promoted_direct_dies() {
let client = test_node("fc/5/client");
let server = test_node("fc/5/server");
let relayed = duplex_establisher(
&client,
&server,
TraversalKind::Relayed,
11,
Duration::ZERO,
None,
);
let kill = Arc::new(Notify::new());
let direct = duplex_establisher(
&client,
&server,
TraversalKind::Direct,
22,
Duration::from_millis(120),
Some(Arc::clone(&kill)),
);
let conn = connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_millis(100),
Duration::from_secs(5),
)
.await
.unwrap();
let mut rx = conn.subscribe();
tokio::time::timeout(Duration::from_secs(5), async {
while *rx.borrow_and_update() != TraversalKind::Direct {
rx.changed().await.unwrap();
}
})
.await
.expect("promoted to Direct");
let post = conn.query_availability(vec![avail_item()]).await.unwrap();
assert_eq!(post.items[0].total_length, Some(22), "served over direct");
kill.notify_waiters();
tokio::time::timeout(Duration::from_secs(5), async {
while *rx.borrow_and_update() != TraversalKind::Relayed {
rx.changed().await.unwrap();
}
})
.await
.expect("fell back to Relayed after direct death");
assert_eq!(conn.current_method(), TraversalKind::Relayed);
let after = conn.query_availability(vec![avail_item()]).await.unwrap();
assert_eq!(after.items[0].total_length, Some(11), "re-dialed relayed");
}
#[tokio::test]
async fn refuses_promotion_to_a_post_tls_blackhole() {
let client = test_node("fc/7/client");
let server = test_node("fc/7/server");
let (relayed, _status) = relayed_establisher(&client, &server, 11);
let direct = blackhole_direct_establisher(&client, &server, Duration::from_millis(100));
let conn = connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_millis(200), Duration::from_millis(300), )
.await
.unwrap();
assert_eq!(conn.current_method(), TraversalKind::Relayed);
tokio::time::sleep(Duration::from_millis(700)).await;
assert_eq!(
conn.current_method(),
TraversalKind::Relayed,
"post-TLS blackhole refused — probe timed out, stays relayed"
);
let resp = conn.query_availability(vec![avail_item()]).await.unwrap();
assert_eq!(resp.items[0].total_length, Some(11), "still relayed");
}
#[tokio::test]
async fn refuses_promotion_on_bls_mismatch_same_peer_id() {
let client = test_node("fc/8/client");
let server = test_node("fc/8/server");
let (relayed, _status) = relayed_establisher(&client, &server, 11);
let direct = mismatched_bls_establisher(&client, &server, 22, Duration::from_millis(100));
let conn = connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_millis(200),
Duration::from_secs(5),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(400)).await;
assert_eq!(
conn.current_method(),
TraversalKind::Relayed,
"promotion refused on BLS mismatch despite matching peer_id"
);
let resp = conn.query_availability(vec![avail_item()]).await.unwrap();
assert_eq!(resp.items[0].total_length, Some(11), "still relayed");
}
#[tokio::test]
async fn inflight_relayed_stream_survives_short_grace_cap() {
let client = test_node("fc/9/client");
let server = test_node("fc/9/server");
let (relayed, _status) = relayed_establisher(&client, &server, 11);
let direct = direct_establisher(&client, &server, 22, Duration::from_millis(150));
let conn = connect_fast_with(
server.peer_id(),
direct,
Some(relayed),
Duration::from_millis(1), Duration::from_secs(5),
)
.await
.expect("relayed lands first");
assert_eq!(conn.current_method(), TraversalKind::Relayed);
let mut pre = conn.open_stream().await.unwrap();
pre.write_all(
&AvailabilityRequest {
items: vec![avail_item()],
}
.encode(),
)
.await
.unwrap();
pre.flush().await.unwrap();
let mut rx = conn.subscribe();
tokio::time::timeout(Duration::from_secs(5), async {
while *rx.borrow_and_update() != TraversalKind::Direct {
rx.changed().await.unwrap();
}
})
.await
.expect("promoted to Direct");
tokio::time::sleep(Duration::from_millis(200)).await;
let pre_resp = AvailabilityResponse::decode(&mut pre).await.unwrap();
assert_eq!(
pre_resp.items[0].total_length,
Some(11),
"in-flight relayed stream survived the short grace cap"
);
}
#[test]
fn fallback_backoff_is_zero_for_a_lone_death_then_capped_exponential() {
let base = Duration::from_millis(50);
let cap = Duration::from_secs(5);
assert_eq!(fallback_backoff(0, base, cap), Duration::ZERO);
assert_eq!(fallback_backoff(1, base, cap), Duration::from_millis(50));
assert_eq!(fallback_backoff(2, base, cap), Duration::from_millis(100));
assert_eq!(fallback_backoff(3, base, cap), Duration::from_millis(200));
assert_eq!(fallback_backoff(20, base, cap), cap);
assert_eq!(fallback_backoff(u32::MAX, base, cap), cap);
}
fn avail_item() -> AvailabilityItem {
AvailabilityItem {
store_id: "bb".repeat(32),
root: None,
retrieval_key: None,
}
}
async fn wait_until(budget: Duration, mut cond: impl FnMut() -> bool) {
let deadline = tokio::time::Instant::now() + budget;
while tokio::time::Instant::now() < deadline {
if cond() {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
}
}