use std::sync::Arc;
use chia_protocol::Bytes32;
use chia_traits::Streamable as _;
use dig_message::envelope::{DigMessageEnvelope, InteractionShape};
use dig_message::{open_message, seal_message, ReplayGuard, SealParams};
use dig_nat::{BindingPolicy, PeerSession, PeerTarget, RangeFrame};
use dig_peer::{DigPeer, NodeCert, SealingIdentity};
use dig_rpc_protocol::envelope::{JsonRpcRequest, JsonRpcResponse};
use dig_rpc_protocol::types::{
FetchModuleRangeParams, GetModuleInfoParams, Health, ModuleInfo, NetworkInfo, RelayStatus,
};
use dig_tls::bls::SecretKey;
use dig_tls::PeerId;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use tokio_rustls::TlsAcceptor;
fn identity_key(label: &str) -> SecretKey {
let mut seed = [0u8; 32];
let bytes = label.as_bytes();
seed[..bytes.len().min(32)].copy_from_slice(&bytes[..bytes.len().min(32)]);
SecretKey::from_seed(&seed)
}
async fn read_framed<R: AsyncReadExt + Unpin>(r: &mut R) -> std::io::Result<Vec<u8>> {
let mut len = [0u8; 4];
r.read_exact(&mut len).await?;
let n = u32::from_be_bytes(len) as usize;
let mut body = vec![0u8; n];
r.read_exact(&mut body).await?;
Ok(body)
}
async fn write_framed<W: AsyncWriteExt + Unpin>(w: &mut W, body: &[u8]) -> std::io::Result<()> {
w.write_all(&(body.len() as u32).to_be_bytes()).await?;
w.write_all(body).await?;
w.flush().await
}
struct TestServer {
addr: std::net::SocketAddr,
peer_id: PeerId,
}
async fn spawn_test_server(server_key: SecretKey) -> TestServer {
let server_node = Arc::new(NodeCert::generate_signed(&server_key).expect("server cert"));
let peer_id = server_node.peer_id();
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let server_tls =
dig_tls::server_config(&server_node, BindingPolicy::Opportunistic).expect("server cfg");
let acceptor = TlsAcceptor::from(server_tls.config.clone());
let (tcp, _) = listener.accept().await.expect("accept tcp");
let tls = acceptor.accept(tcp).await.expect("accept tls");
let client_peer_id = server_tls
.captured_peer_id
.get()
.expect("client peer_id captured");
let client_bls = server_tls.captured_bls.get().expect("client bls captured");
let mut session = PeerSession::server(tls);
let mut counter = 0u64;
while let Some(mut stream) = session.accept_stream().await {
let body = match read_framed(&mut stream).await {
Ok(b) => b,
Err(_) => break,
};
let response = handle_request(
&body,
&server_key,
server_node.peer_id(),
client_peer_id,
&client_bls,
&mut counter,
);
write_framed(&mut stream, &response)
.await
.expect("write response");
}
});
TestServer { addr, peer_id }
}
fn handle_request(
body: &[u8],
server_key: &SecretKey,
server_peer_id: PeerId,
client_peer_id: PeerId,
client_bls: &[u8; 48],
counter: &mut u64,
) -> Vec<u8> {
if let Ok(req) = serde_json::from_slice::<JsonRpcRequest<serde_json::Value>>(body) {
let result = public_result(&req.method);
let response = JsonRpcResponse::success(req.id, result);
return serde_json::to_vec(&response).unwrap();
}
let envelope = DigMessageEnvelope::from_bytes(body).expect("sealed envelope decodes");
let mut guard = ReplayGuard::default();
let resolver = |_did: Bytes32, _epoch: u32| Some(*client_bls);
let opened = open_message(server_key, &envelope, resolver, &mut guard, now_ms())
.expect("server opens the sealed request");
let req: JsonRpcRequest<serde_json::Value> =
serde_json::from_slice(&opened.payload).expect("inner request parses");
let response = JsonRpcResponse::success(req.id, directed_result(&req.method));
let response_json = serde_json::to_vec(&response).unwrap();
*counter += 1;
let params = SealParams {
sender_sk: server_key,
sender: Bytes32::new(*server_peer_id.as_bytes()),
sender_epoch: 0,
recipient: Bytes32::new(*client_peer_id.as_bytes()),
recipient_pub: client_bls,
message_type: dig_peer::seal::RPC_MESSAGE_TYPE,
shape: InteractionShape::Response,
correlation_id: opened.correlation_id,
stream: None,
counter: *counter,
timestamp_ms: now_ms(),
expires_at: 0,
payload: &response_json,
};
seal_message(¶ms)
.expect("server seals the response")
.to_bytes()
.expect("sealed response serializes")
}
fn public_result(method: &str) -> serde_json::Value {
match method {
"dig.methods" => serde_json::json!({ "methods": ["dig.health", "dig.methods"] }),
_ => {
let health = Health {
status: "ok".into(),
version: Some("test".into()),
network_id: Some("DIG_TESTNET".into()),
methods: vec!["dig.health".into()],
};
serde_json::to_value(health).unwrap()
}
}
}
fn directed_result(method: &str) -> serde_json::Value {
match method {
"dig.getPeers" => serde_json::json!({ "peers": [] }),
"dig.announce" => serde_json::json!({ "accepted": true, "known_peers": 1 }),
_ => {
let info = NetworkInfo {
peer_id: None,
network_id: "DIG_TESTNET".into(),
listen_addr: "127.0.0.1:1".into(),
reflexive_addr: None,
candidate_addresses: vec![],
reachability: "direct".into(),
relay: RelayStatus {
url: "off".into(),
reserved: false,
},
};
serde_json::to_value(info).unwrap()
}
}
}
fn now_ms() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_millis() as u64
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn health_round_trips_over_real_mtls() {
let server = spawn_test_server(identity_key("srv/health")).await;
let client_node = Arc::new(NodeCert::generate_signed(&identity_key("cli/health")).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let mut peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect");
assert_eq!(peer.peer_id(), server.peer_id);
let health = peer.health().await.expect("health rpc");
assert_eq!(health.status, "ok");
peer.disconnect().await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn directed_rpc_is_sealed_and_round_trips() {
let server = spawn_test_server(identity_key("srv/net")).await;
let client_key = identity_key("cli/net");
let client_node = Arc::new(NodeCert::generate_signed(&client_key).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let mut peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect")
.with_sealing_identity(SealingIdentity::new(client_key, 0));
assert!(
peer.peer_bls_pub().is_some(),
"peer BLS key must be captured for sealing"
);
assert!(peer.is_sealable());
let info = peer
.get_network_info()
.await
.expect("sealed getNetworkInfo rpc");
assert_eq!(info.network_id, "DIG_TESTNET");
peer.disconnect().await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn directed_rpc_without_sealing_identity_is_refused() {
let server = spawn_test_server(identity_key("srv/refuse")).await;
let client_node = Arc::new(NodeCert::generate_signed(&identity_key("cli/refuse")).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let mut peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect");
let result = peer.get_network_info().await;
assert!(
matches!(result, Err(dig_peer::DigPeerError::NoSealingIdentity)),
"a directed call without a sealing identity must fail closed, got {result:?}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn methods_round_trips() {
let server = spawn_test_server(identity_key("srv/methods")).await;
let client_node = Arc::new(NodeCert::generate_signed(&identity_key("cli/methods")).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let mut peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect");
let methods = peer.methods().await.expect("methods rpc");
assert!(methods.methods.contains(&"dig.health".to_string()));
peer.disconnect().await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn get_peers_and_announce_round_trip_sealed() {
use dig_rpc_protocol::types::AnnounceParams;
let server = spawn_test_server(identity_key("srv/px")).await;
let client_key = identity_key("cli/px");
let client_node = Arc::new(NodeCert::generate_signed(&client_key).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let mut peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect")
.with_sealing_identity(SealingIdentity::new(client_key, 0));
let peers = peer.get_peers().await.expect("sealed getPeers");
assert!(peers.peers.is_empty());
let ack = peer
.announce(&AnnounceParams {
peer_id: peer.peer_id().to_hex(),
addresses: vec![],
})
.await
.expect("sealed announce");
assert!(ack.accepted);
assert_eq!(ack.known_peers, 1);
peer.disconnect().await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn wrong_expected_peer_id_is_rejected() {
let server = spawn_test_server(identity_key("srv/pin")).await;
let client_node = Arc::new(NodeCert::generate_signed(&identity_key("cli/pin")).unwrap());
let wrong_peer_id = PeerId::from_bytes([0xEE; 32]);
let target = PeerTarget::with_addr(wrong_peer_id, server.addr, "DIG_TESTNET");
let result = DigPeer::connect(&target, &client_node).await;
assert!(
result.is_err(),
"connecting with a mismatched peer_id must fail, got Ok"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn open_stream_round_trips_arbitrary_caller_owned_bytes() {
let server = spawn_echo_server(identity_key("srv/raw")).await;
let client_node = Arc::new(NodeCert::generate_signed(&identity_key("cli/raw")).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let mut peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect");
let frame: &[u8] = &[0x00, 0x01, 0xFF, 0xFE, 0x42, 0x00, 0x99];
let mut stream = peer.open_stream().await.expect("open raw stream");
write_framed(&mut stream, frame).await.expect("write frame");
let echoed = read_framed(&mut stream).await.expect("read echoed frame");
assert_eq!(
echoed, frame,
"raw stream must round-trip bytes byte-identically"
);
peer.disconnect().await;
}
async fn spawn_echo_server(server_key: SecretKey) -> TestServer {
let server_node = Arc::new(NodeCert::generate_signed(&server_key).expect("server cert"));
let peer_id = server_node.peer_id();
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let server_tls =
dig_tls::server_config(&server_node, BindingPolicy::Opportunistic).expect("server cfg");
let acceptor = TlsAcceptor::from(server_tls.config.clone());
let (tcp, _) = listener.accept().await.expect("accept tcp");
let tls = acceptor.accept(tcp).await.expect("accept tls");
let mut session = PeerSession::server(tls);
while let Some(mut stream) = session.accept_stream().await {
let body = match read_framed(&mut stream).await {
Ok(b) => b,
Err(_) => break,
};
write_framed(&mut stream, &body).await.expect("echo body");
}
});
TestServer { addr, peer_id }
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn rpc_after_disconnect_is_invalid_state() {
let server = spawn_test_server(identity_key("srv/close")).await;
let client_node = Arc::new(NodeCert::generate_signed(&identity_key("cli/close")).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect");
assert_eq!(peer.state(), dig_peer::PeerState::Connected);
peer.disconnect().await;
}
async fn spawn_module_server(
server_key: SecretKey,
descriptor: ModuleInfo,
blob: Vec<u8>,
frame_size: usize,
) -> TestServer {
let server_node = Arc::new(NodeCert::generate_signed(&server_key).expect("server cert"));
let peer_id = server_node.peer_id();
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let server_tls =
dig_tls::server_config(&server_node, BindingPolicy::Opportunistic).expect("server cfg");
let acceptor = TlsAcceptor::from(server_tls.config.clone());
let (tcp, _) = listener.accept().await.expect("accept tcp");
let tls = acceptor.accept(tcp).await.expect("accept tls");
let mut session = PeerSession::server(tls);
while let Some(mut stream) = session.accept_stream().await {
let Ok(body) = read_framed(&mut stream).await else {
break;
};
let req: JsonRpcRequest<serde_json::Value> =
serde_json::from_slice(&body).expect("module request parses");
match req.method.as_str() {
"dig.getModuleInfo" => {
let response = JsonRpcResponse::success(
req.id,
serde_json::to_value(&descriptor).unwrap(),
);
write_framed(&mut stream, &serde_json::to_vec(&response).unwrap())
.await
.expect("write descriptor");
}
"dig.fetchModuleRange" => {
let params = req.params.expect("module range params");
let offset = params["offset"].as_u64().unwrap_or(0) as usize;
let length = params["length"].as_u64().expect("length") as usize;
let start = offset.min(blob.len());
let window = &blob[start..(start + length).min(blob.len())];
let mut written = 0usize;
while written < window.len() {
let take = frame_size.min(window.len() - written);
let frame = RangeFrame::data(
(start + written) as u64,
window[written..written + take].to_vec(),
)
.with_complete(written + take == window.len());
write_framed(&mut stream, &serde_json::to_vec(&frame).unwrap())
.await
.expect("write module frame");
written += take;
}
}
other => panic!("unexpected module method {other}"),
}
}
});
TestServer { addr, peer_id }
}
fn module_fixture(chunks: usize, chunk_len: usize) -> (ModuleInfo, Vec<u8>) {
let blob: Vec<u8> = (0..chunks * chunk_len).map(|i| (i % 251) as u8).collect();
let info = ModuleInfo {
total_size: blob.len() as u64,
module_hash: sha256_hex(&blob),
chunk_hashes: blob.chunks(chunk_len).map(sha256_hex).collect(),
chunk_lens: vec![chunk_len as u64; chunks],
};
(info, blob)
}
fn sha256_hex(bytes: &[u8]) -> String {
use sha2::{Digest as _, Sha256};
Sha256::digest(bytes)
.iter()
.map(|b| format!("{b:02x}"))
.collect()
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn get_module_info_round_trips_over_real_mtls() {
let (descriptor, blob) = module_fixture(3, 64);
let server =
spawn_module_server(identity_key("srv/modinfo"), descriptor.clone(), blob, 64).await;
let client_node = Arc::new(NodeCert::generate_signed(&identity_key("cli/modinfo")).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let mut peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect");
let info = peer
.get_module_info(&GetModuleInfoParams {
store_id: "aa".repeat(32),
root: "bb".repeat(32),
})
.await
.expect("getModuleInfo rpc");
assert_eq!(info, descriptor, "the descriptor crossed the wire intact");
peer.disconnect().await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn fetch_module_range_streams_the_exact_window() {
let (descriptor, blob) = module_fixture(2, 256);
let server =
spawn_module_server(identity_key("srv/modrange"), descriptor, blob.clone(), 100).await;
let client_node = Arc::new(NodeCert::generate_signed(&identity_key("cli/modrange")).unwrap());
let target = PeerTarget::with_addr(server.peer_id, server.addr, "DIG_TESTNET");
let mut peer = DigPeer::connect(&target, &client_node)
.await
.expect("connect");
let mut stream = peer
.fetch_module_range(&FetchModuleRangeParams {
store_id: "aa".repeat(32),
root: "bb".repeat(32),
offset: Some(256),
length: 256,
})
.await
.expect("fetchModuleRange stream");
let mut got = Vec::new();
loop {
let frame = RangeFrame::decode(&mut stream)
.await
.expect("frame decodes")
.expect("the stream did not end mid-range");
got.extend_from_slice(&frame.bytes);
if frame.complete {
break;
}
}
assert_eq!(got, blob[256..512], "the exact requested window arrived");
peer.disconnect().await;
}