use crate::dht::{
channel::message_channel,
errors::DHTError,
keys::DHTKeyMaterial,
processor::DHTProcessor,
types::{DHTMessageClient, DHTNetworkInfo, DHTRecord, DHTRequest, DHTResponse},
DHTConfig, DefaultRecordValidator, RecordValidator,
};
use libp2p;
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())),
}
};
}
#[derive(Clone, PartialEq, Eq, Debug)]
pub enum DHTStatus {
Initialized,
Active,
Terminated,
Error(String),
}
pub struct DHTNode<V: RecordValidator + 'static> {
config: DHTConfig,
state: DHTStatus,
client: Option<DHTMessageClient>,
thread_handle: Option<tokio::task::JoinHandle<Result<(), DHTError>>>,
keypair: libp2p::identity::Keypair,
peer_id: libp2p::PeerId,
p2p_address: Option<libp2p::Multiaddr>,
bootstrap_peers: Option<Vec<libp2p::Multiaddr>>,
validator: Option<V>,
}
impl<V> DHTNode<V>
where
V: RecordValidator + 'static,
{
pub fn new<K: DHTKeyMaterial>(
key_material: &K,
bootstrap_peers: Option<&Vec<libp2p::Multiaddr>>,
validator: V,
config: &DHTConfig,
) -> Result<Self, DHTError> {
let keypair = key_material.to_dht_keypair()?;
let peer_id = libp2p::PeerId::from(keypair.public());
let peers: Option<Vec<libp2p::Multiaddr>> = bootstrap_peers.map(|peers| peers.to_vec());
let p2p_address: Option<libp2p::Multiaddr> =
if let Some(listening_address) = config.listening_address.as_ref() {
let mut p2p_address = listening_address.to_owned();
p2p_address.push(libp2p::multiaddr::Protocol::P2p(peer_id.into()));
Some(p2p_address)
} else {
None
};
Ok(DHTNode {
keypair,
peer_id,
p2p_address,
config: config.to_owned(),
bootstrap_peers: peers,
state: DHTStatus::Initialized,
client: None,
thread_handle: None,
validator: Some(validator),
})
}
pub fn run(&mut self) -> Result<(), DHTError> {
let (client, processor) = message_channel::<DHTRequest, DHTResponse, DHTError>();
self.ensure_state(DHTStatus::Initialized)?;
self.client = Some(client);
self.thread_handle = Some(DHTProcessor::spawn(
&self.keypair,
&self.peer_id,
&self.p2p_address,
&self.bootstrap_peers,
self.validator.take(),
&self.config,
processor,
)?);
self.state = DHTStatus::Active;
Ok(())
}
pub fn terminate(&mut self) -> Result<(), DHTError> {
self.ensure_state(DHTStatus::Active)?;
if let Some(thread_handle) = self.thread_handle.take() {
thread_handle.abort();
}
self.state = DHTStatus::Terminated;
Ok(())
}
pub fn add_peers(&mut self, new_peers: &[libp2p::Multiaddr]) -> Result<(), DHTError> {
self.ensure_state(DHTStatus::Initialized)?;
let mut new_peers_list: Vec<libp2p::Multiaddr> = new_peers.to_vec();
if let Some(ref mut peers) = self.bootstrap_peers {
peers.append(&mut new_peers_list);
} else {
self.bootstrap_peers = Some(new_peers_list);
}
Ok(())
}
pub fn config(&self) -> &DHTConfig {
&self.config
}
pub fn peer_id(&self) -> &libp2p::PeerId {
&self.peer_id
}
pub fn p2p_address(&self) -> Option<&libp2p::Multiaddr> {
self.p2p_address.as_ref()
}
pub fn status(&self) -> DHTStatus {
self.state.clone()
}
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<DHTNetworkInfo, DHTError> {
let request = DHTRequest::GetNetworkInfo;
let response = self.send_request(request).await?;
ensure_response!(response, DHTResponse::GetNetworkInfo(info) => Ok(info))
}
pub async fn put_record(&self, key: &[u8], value: &[u8]) -> Result<Vec<u8>, DHTError> {
let request = DHTRequest::PutRecord {
key: key.to_vec(),
value: value.to_vec(),
};
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::StartProviding { key: _ } => Ok(()))
}
pub async fn get_providers(&self, key: &[u8]) -> Result<Vec<libp2p::PeerId>, DHTError> {
let request = DHTRequest::GetProviders { key: key.to_vec() };
let response = self.send_request(request).await?;
ensure_response!(response, DHTResponse::GetProviders { providers, key: _ } => Ok(providers))
}
async fn send_request(&self, request: DHTRequest) -> Result<DHTResponse, DHTError> {
self.ensure_state(DHTStatus::Active)?;
self.client
.as_ref()
.expect("active DHT has client")
.send_request_async(request)
.await
.map_err(DHTError::from)
.and_then(|res| res)
}
fn ensure_state(&self, expected_status: DHTStatus) -> Result<(), DHTError> {
if self.state != expected_status {
if expected_status == DHTStatus::Active {
Err(DHTError::NotConnected)
} else {
Err(DHTError::Error("invalid state".into()))
}
} else {
Ok(())
}
}
pub fn validator() -> DefaultRecordValidator {
DefaultRecordValidator {}
}
}
impl<V> Drop for DHTNode<V>
where
V: RecordValidator + 'static,
{
fn drop(&mut self) {
if let Some(thread_handle) = self.thread_handle.take() {
thread_handle.abort();
}
}
}