use std::sync::Arc;
use std::time::Duration;
use sha2::{Digest, Sha256};
use tokio_rustls::TlsAcceptor;
use crate::dialer::MtlsDialer;
use crate::method::relayed::ReservationRelayedTransport;
use crate::method::{MethodOutcome, TraversalKind};
use crate::mux::{AvailabilityAnswer, AvailabilityItem, AvailabilityRequest, AvailabilityResponse};
use crate::peer::PeerTarget;
use crate::relay::loopback_reservation_pair;
use crate::strategy::Dialer;
use crate::tunnel::RelayTunnelStream;
use crate::{BindingPolicy, NodeCert, PeerSession};
use dig_tls::bls::{public_key_bytes, SecretKey};
use tokio::io::AsyncWriteExt;
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"))
}
const RELAY_ENDPOINT: &str = "127.0.0.1:3478";
const NET: &str = "DIG_MAINNET";
#[tokio::test]
async fn relayed_dial_preserves_mtls_peer_id_and_bls_binding() {
let server = test_node("relayed/server");
let client = test_node("relayed/client");
let server_id = server.peer_id();
let client_hex = client.peer_id().to_hex();
let server_hex = server_id.to_hex();
let expected_bls = public_key_bytes(&test_bls_sk("relayed/server"));
let (client_status, server_status) = loopback_reservation_pair(&client_hex, &server_hex);
let server_tunnel = server_status.open_server_tunnel(&client_hex, NET);
let server_tls = dig_tls::server_config(&server, BindingPolicy::Required)
.expect("server config")
.config;
let acceptor = TlsAcceptor::from(server_tls);
tokio::spawn(async move {
let stream = RelayTunnelStream::new(server_tunnel);
let tls = acceptor.accept(stream).await.expect("server mTLS accept");
let mut session = PeerSession::server(tls);
while let Some(mut s) = session.accept_stream().await {
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(77),
chunk_count: Some(1),
complete: Some(true),
})
.collect(),
};
let _ = s.write_all(&resp.encode()).await;
let _ = s.shutdown().await;
}
});
}
});
let transport = Arc::new(ReservationRelayedTransport::new(
Arc::clone(&client_status),
RELAY_ENDPOINT.parse().unwrap(),
));
let dialer = MtlsDialer::new(Arc::clone(&client))
.with_binding_policy(BindingPolicy::Required)
.with_relayed_dialer(transport);
let peer = PeerTarget::relay_only(server_id, NET);
let outcome = MethodOutcome::single(TraversalKind::Relayed, RELAY_ENDPOINT.parse().unwrap());
let mut conn = tokio::time::timeout(Duration::from_secs(5), dialer.dial(&peer, &outcome))
.await
.expect("relayed dial completes")
.expect("relayed dial succeeds");
assert_eq!(conn.peer_id, server_id, "relayed peer_id == server cert id");
assert_eq!(
conn.method,
TraversalKind::Relayed,
"reports the relayed tier"
);
assert_eq!(
conn.peer_bls_pub,
Some(expected_bls),
"relayed dial captured the server's #1204 BLS binding — same as a direct dial"
);
let resp = conn
.query_availability(vec![AvailabilityItem {
store_id: "bb".repeat(32),
root: None,
retrieval_key: None,
}])
.await
.expect("availability over relayed mTLS");
assert_eq!(resp.items.len(), 1);
assert!(resp.items[0].available);
assert_eq!(resp.items[0].total_length, Some(77));
}
#[tokio::test]
async fn relayed_connect_negotiates_client_and_server_roles() {
use crate::RelayAcceptor;
let server = test_node("relayed/role-server");
let client = test_node("relayed/role-client");
let server_id = server.peer_id();
let client_id = client.peer_id();
let client_hex = client_id.to_hex();
let server_hex = server_id.to_hex();
let (client_status, server_status) = loopback_reservation_pair(&client_hex, &server_hex);
let mut inbound = server_status.enable_accept();
let acceptor =
RelayAcceptor::new(Arc::clone(&server)).with_binding_policy(BindingPolicy::Required);
let (served_tx, served_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let tunnel = inbound
.recv()
.await
.expect("introduced circuit surfaces to the acceptor");
let mut conn = acceptor
.accept(tunnel)
.await
.expect("server accepts + completes the mTLS handshake");
let _ = served_tx.send(conn.peer_id);
while let Some(mut s) = conn.session.accept_stream().await {
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(42),
chunk_count: Some(1),
complete: Some(true),
})
.collect(),
};
let _ = s.write_all(&resp.encode()).await;
let _ = s.shutdown().await;
}
});
}
});
let transport = Arc::new(ReservationRelayedTransport::new(
Arc::clone(&client_status),
RELAY_ENDPOINT.parse().unwrap(),
));
let dialer = MtlsDialer::new(Arc::clone(&client))
.with_binding_policy(BindingPolicy::Required)
.with_relayed_dialer(transport);
let peer = PeerTarget::relay_only(server_id, NET);
let outcome = MethodOutcome::single(TraversalKind::Relayed, RELAY_ENDPOINT.parse().unwrap());
let mut conn = tokio::time::timeout(Duration::from_secs(5), dialer.dial(&peer, &outcome))
.await
.expect("relayed dial completes (no double-ClientHello deadlock)")
.expect("relayed dial succeeds");
assert_eq!(
conn.peer_id, server_id,
"client verified the server's peer_id"
);
let served = tokio::time::timeout(Duration::from_secs(5), served_rx)
.await
.expect("server completed its accept")
.expect("server reported the authenticated peer_id");
assert_eq!(served, client_id, "server verified the client's peer_id");
let resp = conn
.query_availability(vec![AvailabilityItem {
store_id: "cc".repeat(32),
root: None,
retrieval_key: None,
}])
.await
.expect("availability round-trips over the role-negotiated relay mTLS");
assert_eq!(resp.items.len(), 1);
assert!(resp.items[0].available);
assert_eq!(resp.items[0].total_length, Some(42));
}
#[tokio::test]
async fn relayed_dial_rejects_malicious_relay_redirect_to_impostor() {
let honest = test_node("relayed/honest"); let impostor = test_node("relayed/impostor"); let client = test_node("relayed/client-b");
let honest_id = honest.peer_id();
let client_hex = client.peer_id().to_hex();
let honest_hex = honest_id.to_hex();
let (client_status, impostor_status) = loopback_reservation_pair(&client_hex, &honest_hex);
let imp_tunnel = impostor_status.open_server_tunnel(&client_hex, NET);
let imp_tls = dig_tls::server_config(&impostor, BindingPolicy::Off)
.unwrap()
.config;
let acceptor = TlsAcceptor::from(imp_tls);
tokio::spawn(async move {
let stream = RelayTunnelStream::new(imp_tunnel);
let _ = acceptor.accept(stream).await; });
let transport = Arc::new(ReservationRelayedTransport::new(
Arc::clone(&client_status),
RELAY_ENDPOINT.parse().unwrap(),
));
let dialer = MtlsDialer::new(Arc::clone(&client)).with_relayed_dialer(transport);
let peer = PeerTarget::relay_only(honest_id, NET);
let outcome = MethodOutcome::single(TraversalKind::Relayed, RELAY_ENDPOINT.parse().unwrap());
let err = tokio::time::timeout(Duration::from_secs(20), dialer.dial(&peer, &outcome))
.await
.expect("relayed dial completes")
.unwrap_err();
assert_eq!(err.kind, TraversalKind::Relayed);
assert!(
err.reason.contains("mtls handshake") || err.reason.contains("peer_id"),
"relayed handshake rejects the impostor substituted by the relay, got: {}",
err.reason
);
}
#[tokio::test]
async fn mutual_simultaneous_relayed_dial_resolves_glare_to_one_client_one_server() {
use crate::peer::PeerConnection;
use crate::relay::RelayTunnel;
use crate::RelayAcceptor;
use tokio::sync::mpsc as tokio_mpsc;
use tokio::sync::oneshot;
let node_a = test_node("glare/a");
let node_b = test_node("glare/b");
let a_id = node_a.peer_id();
let b_id = node_b.peer_id();
let a_hex = a_id.to_hex();
let b_hex = b_id.to_hex();
let (a_status, b_status) = loopback_reservation_pair(&a_hex, &b_hex);
let a_inbound = a_status.enable_accept();
let b_inbound = b_status.enable_accept();
fn spawn_acceptor(
node: Arc<NodeCert>,
mut inbound: tokio_mpsc::Receiver<RelayTunnel>,
) -> (
oneshot::Receiver<crate::PeerId>,
tokio::task::JoinHandle<()>,
) {
let (served_tx, served_rx) = oneshot::channel();
let handle = tokio::spawn(async move {
let acceptor = RelayAcceptor::new(node).with_binding_policy(BindingPolicy::Required);
let Some(tunnel) = inbound.recv().await else {
return;
};
let Ok(mut conn): Result<PeerConnection, _> = acceptor.accept(tunnel).await else {
return;
};
let _ = served_tx.send(conn.peer_id);
while let Some(mut s) = conn.session.accept_stream().await {
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(99),
chunk_count: Some(1),
complete: Some(true),
})
.collect(),
};
let _ = s.write_all(&resp.encode()).await;
let _ = s.shutdown().await;
}
});
}
});
(served_rx, handle)
}
let (a_served_rx, _a_srv) = spawn_acceptor(Arc::clone(&node_a), a_inbound);
let (b_served_rx, _b_srv) = spawn_acceptor(Arc::clone(&node_b), b_inbound);
let a_dialer = MtlsDialer::new(Arc::clone(&node_a))
.with_binding_policy(BindingPolicy::Required)
.with_relayed_dialer(Arc::new(ReservationRelayedTransport::new(
Arc::clone(&a_status),
RELAY_ENDPOINT.parse().unwrap(),
)));
let b_dialer = MtlsDialer::new(Arc::clone(&node_b))
.with_binding_policy(BindingPolicy::Required)
.with_relayed_dialer(Arc::new(ReservationRelayedTransport::new(
Arc::clone(&b_status),
RELAY_ENDPOINT.parse().unwrap(),
)));
let a_peer = PeerTarget::relay_only(b_id, NET); let b_peer = PeerTarget::relay_only(a_id, NET); let outcome = MethodOutcome::single(TraversalKind::Relayed, RELAY_ENDPOINT.parse().unwrap());
let (a_dial, b_dial) = tokio::time::timeout(Duration::from_secs(10), async {
tokio::join!(
a_dialer.dial(&a_peer, &outcome),
b_dialer.dial(&b_peer, &outcome),
)
})
.await
.expect("mutual relayed dial resolves (no glare deadlock)");
let (mut client_conn, expect_client_id, served_rx) = match (a_dial, b_dial) {
(Ok(c), Err(_)) => (c, a_id, b_served_rx), (Err(_), Ok(c)) => (c, b_id, a_served_rx), (Ok(_), Ok(_)) => panic!("glare must yield exactly one client, got two"),
(Err(ea), Err(eb)) => panic!("both relayed dials failed: {ea:?} / {eb:?}"),
};
let expect_server_id = if expect_client_id == a_id { b_id } else { a_id };
assert_eq!(
client_conn.peer_id, expect_server_id,
"client verified the yielding peer's server identity"
);
let served = tokio::time::timeout(Duration::from_secs(5), served_rx)
.await
.expect("server side completed its accept")
.expect("server reported the authenticated client peer_id");
assert_eq!(
served, expect_client_id,
"server verified the winning peer's client identity"
);
let resp = client_conn
.query_availability(vec![AvailabilityItem {
store_id: "dd".repeat(32),
root: None,
retrieval_key: None,
}])
.await
.expect("availability round-trips over the glare-resolved relay mTLS");
assert_eq!(resp.items.len(), 1);
assert!(resp.items[0].available);
assert_eq!(resp.items[0].total_length, Some(99));
}
#[tokio::test]
async fn relay_tunnel_stream_round_trips_bytes() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (a_status, b_status) = loopback_reservation_pair("aa", "bb");
let a_tunnel = a_status.open_tunnel("bb", NET).unwrap();
let b_tunnel = b_status.open_tunnel("aa", NET).unwrap();
let mut a = RelayTunnelStream::new(a_tunnel);
let mut b = RelayTunnelStream::new(b_tunnel);
a.write_all(b"hello world").await.unwrap();
a.flush().await.unwrap();
let mut first = [0u8; 5];
b.read_exact(&mut first).await.unwrap();
assert_eq!(&first, b"hello");
let mut rest = [0u8; 6];
b.read_exact(&mut rest).await.unwrap();
assert_eq!(&rest, b" world");
b.write_all(b"pong").await.unwrap();
b.flush().await.unwrap();
let mut back = [0u8; 4];
a.read_exact(&mut back).await.unwrap();
assert_eq!(&back, b"pong");
}
#[tokio::test]
async fn relayed_dial_without_transport_is_clean_error() {
let dialer = MtlsDialer::new(test_node("relayed/client-c"));
let peer = PeerTarget::relay_only(crate::PeerId::from_bytes([1u8; 32]), NET);
let outcome = MethodOutcome::single(TraversalKind::Relayed, RELAY_ENDPOINT.parse().unwrap());
let err = dialer.dial(&peer, &outcome).await.unwrap_err();
assert_eq!(err.kind, TraversalKind::Relayed);
assert!(err.reason.contains("no relay data-plane"));
}