use async_trait::async_trait;
use crate::error::DhtError;
use crate::routing::Contact;
use crate::wire::{DhtRequest, DhtResponse};
#[async_trait]
pub trait DhtTransport: Send + Sync {
async fn rpc(
&self,
from: &Contact,
peer: &Contact,
request: &DhtRequest,
) -> Result<DhtResponse, DhtError>;
}
#[cfg(test)]
pub(crate) mod memory {
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::Mutex;
use super::*;
pub type Handler = Arc<dyn Fn(DhtRequest) -> DhtResponse + Send + Sync>;
#[derive(Clone, Default)]
pub struct Swarm {
nodes: Arc<Mutex<HashMap<String, Handler>>>,
offline: Arc<Mutex<HashMap<String, ()>>>,
}
impl Swarm {
pub fn new() -> Self {
Swarm::default()
}
pub async fn register(&self, peer_id: String, handler: Handler) {
self.nodes.lock().await.insert(peer_id, handler);
}
pub async fn set_offline(&self, peer_id: &str) {
self.offline.lock().await.insert(peer_id.to_string(), ());
}
pub fn transport(&self) -> MemoryTransport {
MemoryTransport {
swarm: self.clone(),
}
}
}
pub struct MemoryTransport {
swarm: Swarm,
}
#[async_trait]
impl DhtTransport for MemoryTransport {
async fn rpc(
&self,
_from: &Contact,
peer: &Contact,
request: &DhtRequest,
) -> Result<DhtResponse, DhtError> {
if self.swarm.offline.lock().await.contains_key(&peer.peer_id) {
return Err(DhtError::transport(format!("{} is offline", peer.peer_id)));
}
let handler = {
let nodes = self.swarm.nodes.lock().await;
nodes.get(&peer.peer_id).cloned()
};
match handler {
Some(h) => {
let encoded = request.encode();
let mut cur = std::io::Cursor::new(encoded);
let decoded = DhtRequest::decode(&mut cur)
.await
.map_err(DhtError::transport)?;
let resp = h(decoded);
let mut rcur = std::io::Cursor::new(resp.encode());
DhtResponse::decode(&mut rcur)
.await
.map_err(DhtError::transport)
}
None => Err(DhtError::transport(format!("no route to {}", peer.peer_id))),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use dig_nat::PeerId;
fn contact(b: u8) -> Contact {
Contact::new(&PeerId::from_bytes([b; 32]), vec![])
}
#[tokio::test]
async fn dispatches_to_registered_handler() {
let swarm = Swarm::new();
let c = contact(1);
swarm
.register(
c.peer_id.clone(),
Arc::new(|req| match req {
DhtRequest::Ping { nonce } => DhtResponse::Pong { nonce },
_ => DhtResponse::Error {
code: 1,
message: "unexpected".into(),
},
}),
)
.await;
let t = swarm.transport();
let resp = t
.rpc(&contact(0), &c, &DhtRequest::Ping { nonce: 42 })
.await
.unwrap();
assert_eq!(resp, DhtResponse::Pong { nonce: 42 });
}
#[tokio::test]
async fn unrouted_peer_errors() {
let swarm = Swarm::new();
let t = swarm.transport();
let err = t
.rpc(&contact(0), &contact(9), &DhtRequest::Ping { nonce: 1 })
.await;
assert!(matches!(err, Err(DhtError::Transport(_))));
}
#[tokio::test]
async fn offline_peer_errors() {
let swarm = Swarm::new();
let c = contact(2);
swarm
.register(c.peer_id.clone(), Arc::new(|_| DhtResponse::AddProviderOk))
.await;
swarm.set_offline(&c.peer_id).await;
assert!(t_err(&swarm, &c).await);
}
async fn t_err(swarm: &Swarm, c: &Contact) -> bool {
swarm
.transport()
.rpc(&contact(0), c, &DhtRequest::Ping { nonce: 0 })
.await
.is_err()
}
}
}