use crate::dht::{
errors::DhtError,
processor::DhtProcessor,
rpc::{DhtMessageClient, DhtRequest, DhtResponse},
types::{DhtRecord, NetworkInfo, Peer},
DhtConfig, Validator,
};
use libp2p::{identity::Keypair, Multiaddr, PeerId};
use noosphere_common::channel::message_channel;
use std::time::Duration;
use tokio;
macro_rules! ensure_response {
($response:expr, $matcher:pat => $statement:expr) => {
match $response {
$matcher => $statement,
_ => Err(DhtError::Error("Unexpected".into())),
}
};
}
pub struct DhtNode {
config: DhtConfig,
client: DhtMessageClient,
thread_handle: tokio::task::JoinHandle<Result<(), DhtError>>,
peer_id: PeerId,
}
impl DhtNode {
pub fn new<V: Validator + 'static>(
keypair: Keypair,
config: DhtConfig,
validator: Option<V>,
) -> Result<Self, DhtError> {
let peer_id = PeerId::from(keypair.public());
let channels = message_channel::<DhtRequest, DhtResponse, DhtError>();
let thread_handle =
DhtProcessor::spawn(&keypair, peer_id, validator, config.clone(), channels.1)?;
Ok(DhtNode {
peer_id,
config,
client: channels.0,
thread_handle,
})
}
pub fn config(&self) -> &DhtConfig {
&self.config
}
pub fn peer_id(&self) -> &PeerId {
&self.peer_id
}
pub async fn addresses(&self) -> Result<Vec<Multiaddr>, DhtError> {
let request = DhtRequest::GetAddresses { external: false };
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::GetAddresses(addresses) => Ok(addresses))
}
pub async fn external_addresses(&self) -> Result<Vec<Multiaddr>, DhtError> {
let request = DhtRequest::GetAddresses { external: false };
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::GetAddresses(addresses) => Ok(addresses))
}
pub async fn add_peers(&self, peers: Vec<Multiaddr>) -> Result<(), DhtError> {
let request = DhtRequest::AddPeers { peers };
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::Success => Ok(()))
}
pub async fn listen(&self, listening_address: Multiaddr) -> Result<Multiaddr, DhtError> {
let request = DhtRequest::StartListening {
address: listening_address,
};
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::Address(addr) => Ok(addr))
}
pub async fn stop_listening(&self) -> Result<(), DhtError> {
let request = DhtRequest::StopListening;
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::Success => Ok(()))
}
pub async fn wait_for_peers(&self, requested_peers: usize) -> Result<(), DhtError> {
loop {
let info = self.network_info().await?;
if info.num_peers >= requested_peers {
return Ok(());
}
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
pub async fn bootstrap(&self) -> Result<(), DhtError> {
let request = DhtRequest::Bootstrap;
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::Success => Ok(()))
}
pub async fn network_info(&self) -> Result<NetworkInfo, DhtError> {
let request = DhtRequest::GetNetworkInfo;
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::GetNetworkInfo(info) => Ok(info))
}
pub async fn peers(&self) -> Result<Vec<Peer>, DhtError> {
let request = DhtRequest::GetPeers;
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::GetPeers(peers) => Ok(peers))
}
pub async fn put_record(
&self,
key: &[u8],
value: &[u8],
quorum: usize,
) -> Result<Vec<u8>, DhtError> {
let request = DhtRequest::PutRecord {
key: key.to_vec(),
value: value.to_vec(),
quorum,
};
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::PutRecord { key } => Ok(key))
}
pub async fn get_record(&self, key: &[u8]) -> Result<DhtRecord, DhtError> {
let request = DhtRequest::GetRecord { key: key.to_vec() };
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::GetRecord(record) => Ok(record))
}
pub async fn start_providing(&self, key: &[u8]) -> Result<(), DhtError> {
let request = DhtRequest::StartProviding { key: key.to_vec() };
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::Success => Ok(()))
}
pub async fn get_providers(&self, key: &[u8]) -> Result<Vec<PeerId>, DhtError> {
let request = DhtRequest::GetProviders { key: key.to_vec() };
let response = self.send_request(request).await?;
ensure_response!(response, DhtResponse::GetProviders { providers } => Ok(providers))
}
async fn send_request(&self, request: DhtRequest) -> Result<DhtResponse, DhtError> {
self.client
.send(request)
.await
.map_err(DhtError::from)
.and_then(|res| res)
}
}
impl Drop for DhtNode {
fn drop(&mut self) {
self.thread_handle.abort();
}
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg(test)]
mod test {
use super::*;
use std::fmt::Display;
use crate::dht::{AllowAllValidator, DhtError, DhtNode, NetworkInfo, Validator};
use async_trait::async_trait;
use crate::utils::make_p2p_address;
use futures::future::try_join_all;
use libp2p::{self, Multiaddr};
use std::future::Future;
use std::time::Duration;
pub async fn wait_ms(ms: u64) {
tokio::time::sleep(Duration::from_millis(ms)).await;
}
async fn await_or_timeout<T>(
timeout_ms: u64,
future: impl Future<Output = T>,
message: String,
) -> T {
tokio::select! {
_ = wait_ms(timeout_ms) => { panic!("timed out: {}", message); }
result = future => { result }
}
}
pub async fn swarm_command<'a, TFuture, F, T, E>(
nodes: &'a mut [DhtNode],
func: F,
) -> Result<Vec<T>, E>
where
F: FnMut(&'a mut DhtNode) -> TFuture,
TFuture: Future<Output = Result<T, E>>,
{
let futures: Vec<_> = nodes.iter_mut().map(func).collect();
try_join_all(futures).await
}
async fn create_network<V: Validator + Clone + 'static>(
node_count: usize,
validator: Option<V>,
) -> Result<Vec<DhtNode>, anyhow::Error> {
let mut bootstrap_addresses: Option<Vec<Multiaddr>> = None;
let mut nodes = vec![];
for _ in 0..node_count {
let node = DhtNode::new(
Keypair::generate_ed25519(),
Default::default(),
validator.clone(),
)?;
if let Some(addresses) = bootstrap_addresses.as_ref() {
node.add_peers(addresses.to_owned()).await?;
node.listen("/ip4/127.0.0.1/tcp/0".parse().unwrap()).await?;
} else {
let address = node.listen("/ip4/127.0.0.1/tcp/0".parse().unwrap()).await?;
bootstrap_addresses = Some(vec![address]);
}
nodes.push(node);
}
Ok(nodes)
}
async fn initialize_network(nodes: &mut Vec<DhtNode>) -> Result<(), anyhow::Error> {
let expected_peers = nodes.len() - 1;
wait_ms(700).await;
swarm_command(nodes, |c| c.bootstrap()).await?;
await_or_timeout(
5000,
swarm_command(nodes, |c| c.wait_for_peers(expected_peers)),
format!("waiting for {} peers", expected_peers),
)
.await?;
Ok(())
}
fn create_unfiltered_dht_node() -> Result<DhtNode, DhtError> {
DhtNode::new::<AllowAllValidator>(
Keypair::generate_ed25519(),
Default::default(),
Some(AllowAllValidator {}),
)
}
#[tokio::test]
async fn test_dhtnode_base_case() -> Result<(), DhtError> {
let node = create_unfiltered_dht_node()?;
node.listen("/ip4/127.0.0.1/tcp/0".parse().unwrap()).await?;
let info = node.network_info().await?;
assert_eq!(
info,
NetworkInfo {
num_connections: 0,
num_established: 0,
num_peers: 0,
num_pending: 0,
}
);
if node.bootstrap().await.is_err() {
panic!("bootstrap() should succeed, even without peers to bootstrap.");
}
Ok(())
}
#[tokio::test]
async fn test_dhtnode_bootstrap() -> Result<(), DhtError> {
let num_nodes = 5;
let mut nodes = create_network(num_nodes, Some(AllowAllValidator {})).await?;
initialize_network(&mut nodes).await?;
for info in swarm_command(&mut nodes, |c| c.network_info()).await? {
assert_eq!(info.num_peers, num_nodes - 1);
assert_eq!(info.num_pending, 0);
}
let info = nodes.first().unwrap().network_info().await?;
assert_eq!(info.num_peers, num_nodes - 1);
assert_eq!(info.num_pending, 0);
Ok(())
}
#[tokio::test]
async fn test_dhtnode_simple() -> Result<(), DhtError> {
let mut nodes = create_network(2, Some(AllowAllValidator {})).await?;
initialize_network(&mut nodes).await?;
let (node_a, node_b) = (nodes.pop().unwrap(), nodes.pop().unwrap());
node_a.put_record(b"foo", b"bar", 1).await?;
let result = node_b.get_record(b"foo").await?;
assert_eq!(result.key, b"foo");
assert_eq!(result.value.expect("has value"), b"bar");
Ok(())
}
#[tokio::test]
async fn test_dhtnode_providers() -> Result<(), DhtError> {
let mut nodes = create_network(2, Some(AllowAllValidator {})).await?;
initialize_network(&mut nodes).await?;
let (node_a, node_b) = (nodes.pop().unwrap(), nodes.pop().unwrap());
node_a.start_providing(b"foo").await?;
let providers = node_b.get_providers(b"foo").await?;
assert_eq!(providers.len(), 1);
assert_eq!(&providers[0], node_a.peer_id());
Ok(())
}
#[tokio::test]
async fn test_dhtnode_validator() -> Result<(), DhtError> {
#[derive(Clone)]
struct MyValidator {}
#[async_trait]
impl Validator for MyValidator {
async fn validate(&mut self, data: &[u8]) -> bool {
data == b"VALID"
}
}
impl Display for MyValidator {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "MyValidator")
}
}
let mut nodes = create_network(2, Some(MyValidator {})).await?;
initialize_network(&mut nodes).await?;
let (node_a, node_b) = (nodes.pop().unwrap(), nodes.pop().unwrap());
let unfiltered_client = create_unfiltered_dht_node()?;
unfiltered_client
.add_peers(vec![make_p2p_address(
node_a.addresses().await?.pop().unwrap(),
node_a.peer_id().to_owned(),
)])
.await?;
node_a.put_record(b"foo_1", b"VALID", 1).await?;
let result = node_b.get_record(b"foo_1").await?;
assert_eq!(
result.value.expect("has value"),
b"VALID",
"validation allows valid records through"
);
assert!(
node_a.put_record(b"foo_2", b"INVALID", 1).await.is_err(),
"setting a record validates locally"
);
unfiltered_client.put_record(b"foo_3", b"VALID", 1).await?;
unfiltered_client
.put_record(b"foo_4", b"INVALID", 1)
.await?;
let result = node_b.get_record(b"foo_3").await?;
assert_eq!(
result.value.expect("has value"),
b"VALID",
"validation allows valid records through"
);
assert!(
node_b.get_record(b"foo_4").await?.value.is_none(),
"invalid records are not retrieved from the network"
);
Ok(())
}
}