use super::super::bencode::{Dict, Value, bytes, entry};
use super::super::krpc::{Node, NodeId};
use super::super::mutable::Mutable;
use std::net::{SocketAddr, UdpSocket};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, PoisonError};
use std::thread::JoinHandle;
use std::time::Duration;
use wire::{claim, error, reply, routing};
mod wire;
pub(crate) const TOKEN: &[u8] = b"tok";
#[derive(Clone, Copy)]
pub(crate) enum Mood {
Answer,
Silent,
Garbage,
Refuse,
Anonymous,
Stray,
Router,
Rotor,
Mute,
Claim(SocketAddr),
}
pub(crate) struct FakeNode {
socket: Option<UdpSocket>,
pub(crate) addr: SocketAddr,
pub(crate) id: NodeId,
stop: Arc<AtomicBool>,
thread: Option<JoinHandle<()>>,
items: Arc<Mutex<Vec<Mutable>>>,
queries: Arc<AtomicUsize>,
}
impl FakeNode {
pub(crate) fn bind(id: NodeId) -> FakeNode {
let socket = UdpSocket::bind("127.0.0.1:0").unwrap();
socket
.set_read_timeout(Some(Duration::from_millis(20)))
.unwrap();
let addr = socket.local_addr().unwrap();
FakeNode {
socket: Some(socket),
addr,
id,
stop: Arc::new(AtomicBool::new(false)),
thread: None,
items: Arc::new(Mutex::new(Vec::new())),
queries: Arc::new(AtomicUsize::new(0)),
}
}
pub(crate) fn items(&self) -> Vec<Mutable> {
self.items
.lock()
.unwrap_or_else(PoisonError::into_inner)
.clone()
}
pub(crate) fn store(&self) -> Arc<Mutex<Vec<Mutable>>> {
Arc::clone(&self.items)
}
pub(crate) fn queries(&self) -> usize {
self.queries.load(Ordering::Relaxed)
}
pub(crate) fn node(&self) -> Node {
Node {
id: self.id,
addr: self.addr,
}
}
pub(crate) fn serve(&mut self, peers: Vec<Node>, mood: Mood, items: Vec<Mutable>) {
let socket = self.socket.take().unwrap();
let (id, stop) = (self.id, Arc::clone(&self.stop));
*self.items.lock().unwrap_or_else(PoisonError::into_inner) = items;
let (store, queries) = (Arc::clone(&self.items), Arc::clone(&self.queries));
self.thread = Some(std::thread::spawn(move || {
run(&socket, id, &peers, mood, &store, &queries, &stop);
}));
}
}
impl Drop for FakeNode {
fn drop(&mut self) {
self.stop.store(true, Ordering::Relaxed);
if let Some(t) = self.thread.take() {
t.join().unwrap();
}
}
}
fn run(
socket: &UdpSocket,
id: NodeId,
peers: &[Node],
mood: Mood,
store: &Mutex<Vec<Mutable>>,
queries: &AtomicUsize,
stop: &AtomicBool,
) {
let mut buf = vec![0u8; 8192];
let mut turn = 0usize;
while !stop.load(Ordering::Relaxed) {
let Ok((n, from)) = socket.recv_from(&mut buf) else {
continue;
};
queries.fetch_add(1, Ordering::Relaxed);
let mut items = store.lock().unwrap_or_else(PoisonError::into_inner);
let q = Value::decode(&buf[..n]).unwrap();
let tid = q.get("t").unwrap().as_bytes().unwrap().to_vec();
let verb = q.get("q").unwrap().as_bytes().unwrap();
let datagram = match mood {
Mood::Silent => continue,
Mood::Router | Mood::Rotor if verb != b"find_node" => continue,
Mood::Mute if verb == b"put" => continue,
Mood::Rotor => {
turn += 1;
let one = vec![peers[(turn - 1) % peers.len()]; 8];
answer(&tid, id, &one, &mut items, &q)
}
Mood::Garbage => b"not bencode".to_vec(),
Mood::Refuse => error(&tid, 201, "refused"),
Mood::Anonymous => reply(&tid, Dict::new()),
Mood::Stray => reply(b"stray", Dict::from([entry("id", bytes(&id.0))])),
Mood::Answer | Mood::Router | Mood::Mute => answer(&tid, id, peers, &mut items, &q),
Mood::Claim(ip) => claim(answer(&tid, id, peers, &mut items, &q), ip),
};
socket.send_to(&datagram, from).unwrap();
}
}
fn answer(tid: &[u8], id: NodeId, peers: &[Node], items: &mut Vec<Mutable>, q: &Value) -> Vec<u8> {
let a = q.get("a").unwrap();
let mut r = Dict::from([entry("id", bytes(&id.0))]);
match q.get("q").unwrap().as_bytes().unwrap() {
b"find_node" => {
r.extend(routing(peers));
reply(tid, r)
}
b"get" => {
r.extend(routing(peers));
r.insert(b"token".to_vec(), bytes(TOKEN));
let target = NodeId::parse(a.get("target").unwrap().as_bytes().unwrap()).unwrap();
if let Some(item) = items.iter().find(|i| i.target() == target) {
r.insert(b"k".to_vec(), bytes(&item.key));
r.insert(b"seq".to_vec(), Value::Int(item.seq));
r.insert(b"sig".to_vec(), bytes(&item.sig));
r.insert(b"v".to_vec(), bytes(&item.value));
}
reply(tid, r)
}
b"put" => {
if a.get("token").unwrap().as_bytes().unwrap() != TOKEN {
return error(tid, 203, "bad token");
}
let item = Mutable {
key: a.get("k").unwrap().as_bytes().unwrap().try_into().unwrap(),
salt: a
.get("salt")
.map(|s| s.as_bytes().unwrap().to_vec())
.unwrap_or_default(),
seq: a.get("seq").unwrap().as_int().unwrap(),
value: a.get("v").unwrap().as_bytes().unwrap().to_vec(),
sig: a
.get("sig")
.unwrap()
.as_bytes()
.unwrap()
.try_into()
.unwrap(),
};
if !item.verify() {
return error(tid, 206, "invalid signature");
}
let target = item.target();
if items
.iter()
.any(|i| i.target() == target && i.seq >= item.seq)
{
return error(tid, 302, "sequence number less than current");
}
items.retain(|i| i.target() != target);
items.push(item);
reply(tid, r)
}
_ => error(tid, 204, "unknown method"),
}
}