use async_std::sync::Arc; use std::convert::TryFrom; use std::io; use std::net::SocketAddr; use crate::serial::{deserialize, serialize}; use crate::{Decodable, Encodable, Result}; use async_executor::Executor; use bytes::Bytes; use futures::FutureExt; use log::*; use rand::Rng; use signal_hook::{consts::SIGINT, iterator::Signals}; use zeromq::*; pub type PeerId = Vec; enum NetEvent { Receive(zeromq::ZmqMessage), Send((PeerId, Reply)), Stop, } pub fn addr_to_string(addr: SocketAddr) -> String { format!("tcp://{}", addr.to_string()) } pub struct RepProtocol { addr: SocketAddr, socket: zeromq::RouterSocket, recv_queue: async_channel::Receiver<(PeerId, Reply)>, send_queue: async_channel::Sender<(PeerId, Request)>, channels: ( async_channel::Sender<(PeerId, Reply)>, async_channel::Receiver<(PeerId, Request)>, ), service_name: String, } impl RepProtocol { pub fn new(addr: SocketAddr, service_name: String) -> RepProtocol { let socket = zeromq::RouterSocket::new(); let (send_queue, recv_channel) = async_channel::unbounded::<(PeerId, Request)>(); let (send_channel, recv_queue) = async_channel::unbounded::<(PeerId, Reply)>(); let channels = (send_channel.clone(), recv_channel.clone()); RepProtocol { addr, socket, recv_queue, send_queue, channels, service_name, } } pub async fn start( &mut self, ) -> Result<( async_channel::Sender<(PeerId, Reply)>, async_channel::Receiver<(PeerId, Request)>, )> { let addr = addr_to_string(self.addr); self.socket.bind(addr.as_str()).await?; info!("{} SERVICE: Bound To {}", self.service_name, addr); Ok(self.channels.clone()) } pub async fn run(&mut self, executor: Arc>) -> Result<()> { info!("{} SERVICE: Running", self.service_name); let (stop_s, stop_r) = async_channel::unbounded::<()>(); let mut signals = Signals::new(&[SIGINT])?; let stop_task = executor.spawn(async move { for _ in signals.forever() { stop_s.send(()).await?; break; } Ok::<(), crate::Error>(()) }); loop { let event = futures::select! { msg = self.socket.recv().fuse() => NetEvent::Receive(msg?), msg = self.recv_queue.recv().fuse() => NetEvent::Send(msg?), _ = stop_r.recv().fuse() => NetEvent::Stop }; match event { NetEvent::Receive(msg) => { if let Some(peer) = msg.get(0) { if let Some(request) = msg.get(1) { let request: Vec = request.to_vec(); let request: Request = deserialize(&request)?; self.send_queue.send((peer.to_vec(), request)).await?; } } } NetEvent::Send((peer, reply)) => { let peer = Bytes::from(peer); let mut msg: Vec = vec![peer]; let reply: Vec = serialize(&reply); let reply = Bytes::from(reply); msg.push(reply); let reply = zeromq::ZmqMessage::try_from(msg) .map_err(|_| crate::Error::TryFromError)?; self.socket.send(reply).await?; } NetEvent::Stop => break, } } let _ = stop_task.cancel().await; warn!("{} SERVICE: Stopped", self.service_name); Ok(()) } } pub struct ReqProtocol { addr: SocketAddr, socket: zeromq::DealerSocket, service_name: String, } impl ReqProtocol { pub fn new(addr: SocketAddr, service_name: String) -> ReqProtocol { let socket = zeromq::DealerSocket::new(); ReqProtocol { addr, socket, service_name, } } pub async fn start(&mut self) -> Result<()> { let addr = addr_to_string(self.addr); self.socket.connect(addr.as_str()).await?; info!("{} SERVICE: Connected To {}", self.service_name, self.addr); Ok(()) } pub async fn request( &mut self, command: u8, data: Vec, handle_error: &dyn Fn(u32), ) -> Result>> { let request = Request::new(command, data); let req = serialize(&request); let req = bytes::Bytes::from(req); let req: zeromq::ZmqMessage = req.into(); self.socket.send(req).await?; info!( "{} SERVICE: Sent Request {{ command: {} }}", self.service_name, command ); let rep: zeromq::ZmqMessage = self.socket.recv().await?; if let Some(reply) = rep.get(0) { let reply: Vec = reply.to_vec(); let reply: Reply = deserialize(&reply)?; info!( "{} SERVICE: Received Reply {{ error: {} }}", self.service_name, reply.has_error() ); if reply.has_error() { // TODO return error status code instead of None // this is temporary handle_error(reply.get_error()); return Ok(None); } assert!(reply.get_id() == request.get_id()); Ok(Some(reply.get_payload())) } else { Err(crate::Error::ZmqError( "Couldn't parse ZmqMessage".to_string(), )) } } } pub struct Publisher { addr: SocketAddr, socket: zeromq::PubSocket, service_name: String, } impl Publisher { pub fn new(addr: SocketAddr, service_name: String) -> Publisher { let socket = zeromq::PubSocket::new(); Publisher { addr, socket, service_name, } } pub async fn start(&mut self, recv_queue: async_channel::Receiver>) -> Result<()> { let addr = addr_to_string(self.addr); self.socket.bind(addr.as_str()).await?; info!( "{} PUBLISHER SERVICE : Bound To {}", self.service_name, addr ); loop { let msg = recv_queue.recv().await?; self.publish(msg).await?; } } async fn publish(&mut self, data: Vec) -> Result<()> { let data = Bytes::from(data); self.socket.send(data.into()).await?; Ok(()) } } pub struct Subscriber { addr: SocketAddr, socket: zeromq::SubSocket, service_name: String, } impl Subscriber { pub fn new(addr: SocketAddr, service_name: String) -> Subscriber { let socket = zeromq::SubSocket::new(); Subscriber { addr, socket, service_name, } } pub async fn start(&mut self) -> Result<()> { let addr = addr_to_string(self.addr); self.socket.connect(addr.as_str()).await?; self.socket.subscribe("").await?; info!( "{} SUBSCRIBER SERVICE : Connected To {}", self.service_name, addr ); Ok(()) } pub async fn fetch(&mut self) -> Result { let data = self.socket.recv().await?; match data.get(0) { Some(d) => { let data = d.to_vec(); let data: T = deserialize(&data)?; Ok(data) } None => Err(crate::Error::ZmqError( "Couldn't parse ZmqMessage".to_string(), )), } } } #[derive(Debug, PartialEq)] pub struct Request { command: u8, id: u32, payload: Vec, } impl Request { pub fn new(command: u8, payload: Vec) -> Request { let id = Self::gen_id(); Request { command, id, payload, } } fn gen_id() -> u32 { let mut rng = rand::thread_rng(); rng.gen() } pub fn get_id(&self) -> u32 { self.id } pub fn get_command(&self) -> u8 { self.command } pub fn get_payload(&self) -> Vec { self.payload.clone() } } #[derive(Debug, PartialEq)] pub struct Reply { id: u32, error: u32, payload: Vec, } impl Reply { pub fn from(request: &Request, error: u32, payload: Vec) -> Reply { Reply { id: request.get_id(), error, payload, } } pub fn has_error(&self) -> bool { if self.error == 0 { false } else { true } } pub fn get_error(&self) -> u32 { self.error } pub fn get_payload(&self) -> Vec { self.payload.clone() } pub fn set_payload(&mut self, payload: Vec) { self.payload = payload; } pub fn set_error(&mut self, error: u32) { self.error = error; } pub fn get_id(&self) -> u32 { self.id } } impl Encodable for Request { fn encode(&self, mut s: S) -> Result { let mut len = 0; len += self.command.encode(&mut s)?; len += self.id.encode(&mut s)?; len += self.payload.encode(&mut s)?; Ok(len) } } impl Encodable for Reply { fn encode(&self, mut s: S) -> Result { let mut len = 0; len += self.id.encode(&mut s)?; len += self.error.encode(&mut s)?; len += self.payload.encode(&mut s)?; Ok(len) } } impl Decodable for Request { fn decode(mut d: D) -> Result { Ok(Self { command: Decodable::decode(&mut d)?, id: Decodable::decode(&mut d)?, payload: Decodable::decode(&mut d)?, }) } } impl Decodable for Reply { fn decode(mut d: D) -> Result { Ok(Self { id: Decodable::decode(&mut d)?, error: Decodable::decode(&mut d)?, payload: Decodable::decode(&mut d)?, }) } } #[cfg(test)] mod tests { use super::{Reply, Request, Result}; use crate::serial::{deserialize, serialize}; #[test] fn serialize_and_deserialize_request_test() { let request = Request::new(2, vec![2, 3, 4, 6, 4]); let serialized_request = serialize(&request); assert!((deserialize(&serialized_request) as Result).is_err()); let deserialized_request = deserialize(&serialized_request).ok(); assert_eq!(deserialized_request, Some(request)); } #[test] fn serialize_and_deserialize_reply_test() { let request = Request::new(2, vec![2, 3, 4, 6, 4]); let reply = Reply::from(&request, 0, vec![2, 3, 4, 6, 4]); let serialized_reply = serialize(&reply); assert!((deserialize(&serialized_reply) as Result).is_err()); let deserialized_reply = deserialize(&serialized_reply).ok(); assert_eq!(deserialized_reply, Some(reply)); } }