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_tunnel(&client_hex, NET)
.expect("server opens relay tunnel");
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_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_tunnel(&client_hex, NET).unwrap();
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 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"));
}