use std::collections::HashMap;
use std::io::{Error, ErrorKind, Result};
use std::sync::{Arc, Mutex};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use tokio::sync::{mpsc, oneshot};
use crate::frame::{self, Descriptor, Entry, Frame};
use crate::id;
struct Peer {
tx: mpsc::UnboundedSender<Vec<u8>>,
desc: Descriptor,
token: u64,
evict: oneshot::Sender<()>,
}
#[derive(Clone)]
pub struct Registry {
peers: Arc<Mutex<HashMap<String, Peer>>>,
host: Arc<str>,
}
impl Registry {
pub fn new(host: &str) -> Registry {
Registry {
peers: Arc::default(),
host: host.into(),
}
}
}
static NEXT_TOKEN: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
pub async fn handle<S>(stream: S, registry: Registry, take_over: bool) -> Result<()>
where
S: AsyncRead + AsyncWrite + Send + 'static,
{
let (mut rd, mut wr) = tokio::io::split(stream);
let hello = match Frame::read(&mut rd).await? {
Some(f) if f.typ == frame::HELLO => f,
Some(_) => return Err(Error::new(ErrorKind::InvalidData, "expected HELLO")),
None => return Ok(()),
};
let Some(id) = id::stamp(&hello.from, ®istry.host) else {
let why = format!("bad_id:{}", hello.from);
wr.write_all(&Frame::new(frame::ERROR, "zapd", "", why.clone().into_bytes()).encode())
.await?;
return Err(Error::new(ErrorKind::InvalidData, why));
};
let desc = frame::decode_hello(&hello.payload)?;
let (role, brand) = (desc.role, desc.brand.clone());
let (tx, mut rx) = mpsc::unbounded_channel::<Vec<u8>>();
let (evict, mut evicted_rx) = oneshot::channel();
let token = NEXT_TOKEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let registered = {
let mut peers = registry.peers.lock().unwrap();
if !take_over && peers.contains_key(&id) {
None
} else {
Some(peers.insert(
id.clone(),
Peer {
tx: tx.clone(),
desc,
token,
evict,
},
))
}
};
let Some(evicted) = registered else {
let why = format!("taken:{id}");
wr.write_all(&Frame::new(frame::ERROR, "zapd", "", why.clone().into_bytes()).encode())
.await?;
return Err(Error::new(ErrorKind::AddrInUse, why));
};
if let Some(old) = evicted {
let _ = old.evict.send(());
tracing::info!("zapd: {id} reconnected — replaced stale peer");
}
tracing::info!("zapd: {id} online (role={role}, brand={brand})");
let writer = tokio::spawn(async move {
while let Some(bytes) = rx.recv().await {
if wr.write_all(&bytes).await.is_err() {
break;
}
}
});
let _ = tx.send(Frame::new(frame::WELCOME, "zapd", &id, Vec::new()).encode());
broadcast(
®istry,
&id,
frame::PEER_CONNECTED,
frame::encode_peer(&id),
);
let result = tokio::select! {
r = route_loop(&mut rd, ®istry, &id) => r,
_ = &mut evicted_rx => Ok(()),
};
let was_current = {
let mut reg = registry.peers.lock().unwrap();
if reg.get(&id).map(|p| p.token) == Some(token) {
reg.remove(&id);
true
} else {
false
}
};
if was_current {
broadcast(
®istry,
&id,
frame::PEER_DISCONNECTED,
frame::encode_peer(&id),
);
tracing::info!("zapd: {id} offline");
}
writer.abort();
result
}
async fn route_loop<R: AsyncRead + Unpin>(rd: &mut R, registry: &Registry, id: &str) -> Result<()> {
while let Some(mut f) = Frame::read(rd).await? {
if f.to.is_empty() {
match f.typ {
frame::PROVIDERS_LIST => {
let filter = frame::decode_brand_filter(&f.payload);
let entries = list(registry, &filter);
let reply = Frame::new(
frame::PROVIDERS,
"zapd",
id,
frame::encode_providers(&entries),
);
send_to(registry, id, reply.encode());
}
frame::HELLO => { }
_ => {
let err = Frame::new(frame::ERROR, "zapd", id, b"unknown_control".to_vec());
send_to(registry, id, err.encode());
}
}
} else {
let dest = f.to.clone();
f.from = id.to_string();
if !send_to(registry, &dest, f.encode()) {
let err = Frame::new(
frame::ERROR,
"zapd",
id,
format!("no_route:{dest}").into_bytes(),
);
send_to(registry, id, err.encode());
}
}
}
Ok(())
}
fn list(registry: &Registry, brand_filter: &str) -> Vec<Entry> {
let reg = registry.peers.lock().unwrap();
reg.iter()
.filter(|(_, p)| brand_filter.is_empty() || p.desc.brand == brand_filter)
.map(|(id, p)| Entry {
id: id.clone(),
desc: p.desc.clone(),
})
.collect()
}
fn send_to(registry: &Registry, id: &str, bytes: Vec<u8>) -> bool {
let reg = registry.peers.lock().unwrap();
reg.get(id)
.map(|p| p.tx.send(bytes).is_ok())
.unwrap_or(false)
}
fn broadcast(registry: &Registry, except: &str, typ: u8, payload: Vec<u8>) {
let reg = registry.peers.lock().unwrap();
for (pid, p) in reg.iter() {
if pid == except {
continue;
}
let f = Frame::new(typ, "zapd", pid, payload.clone());
let _ = p.tx.send(f.encode());
}
}