reqrep.rs 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394
  1. use std::{io, net::SocketAddr, sync::Arc};
  2. use async_executor::Executor;
  3. use async_std::prelude::*;
  4. use bytes::Bytes;
  5. use futures::FutureExt;
  6. use log::*;
  7. use rand::Rng;
  8. use signal_hook::consts::SIGINT;
  9. use signal_hook_async_std::Signals;
  10. use zeromq::*;
  11. use crate::{
  12. serial::{deserialize, serialize, Decodable, Encodable},
  13. Result,
  14. };
  15. pub type PeerId = Vec<u8>;
  16. pub type Channels =
  17. (async_channel::Sender<(PeerId, Reply)>, async_channel::Receiver<(PeerId, Request)>);
  18. enum NetEvent {
  19. Receive(zeromq::ZmqMessage),
  20. Send((PeerId, Reply)),
  21. Stop,
  22. }
  23. pub fn addr_to_string(addr: SocketAddr) -> String {
  24. format!("tcp://{}", addr.to_string())
  25. }
  26. pub struct RepProtocol {
  27. addr: SocketAddr,
  28. socket: zeromq::RouterSocket,
  29. recv_queue: async_channel::Receiver<(PeerId, Reply)>,
  30. send_queue: async_channel::Sender<(PeerId, Request)>,
  31. channels: Channels,
  32. service_name: String,
  33. }
  34. impl RepProtocol {
  35. pub fn new(addr: SocketAddr, service_name: String) -> RepProtocol {
  36. let socket = zeromq::RouterSocket::new();
  37. let (send_queue, recv_channel) = async_channel::unbounded::<(PeerId, Request)>();
  38. let (send_channel, recv_queue) = async_channel::unbounded::<(PeerId, Reply)>();
  39. let channels = (send_channel, recv_channel);
  40. RepProtocol { addr, socket, recv_queue, send_queue, channels, service_name }
  41. }
  42. pub async fn start(
  43. &mut self,
  44. ) -> Result<(async_channel::Sender<(PeerId, Reply)>, async_channel::Receiver<(PeerId, Request)>)>
  45. {
  46. let addr = addr_to_string(self.addr);
  47. self.socket.bind(addr.as_str()).await?;
  48. debug!(target: "REP PROTOCOL API", "{} SERVICE: Bound To {}", self.service_name, addr);
  49. Ok(self.channels.clone())
  50. }
  51. pub async fn run(&mut self, executor: Arc<Executor<'_>>) -> Result<()> {
  52. debug!(target: "REP PROTOCOL API", "{} SERVICE: Running", self.service_name);
  53. let (stop_s, stop_r) = async_channel::unbounded::<()>();
  54. let signals = Signals::new(&[SIGINT])?;
  55. let handle = signals.handle();
  56. let signals_task = executor.spawn(async move {
  57. let mut signals = signals.fuse();
  58. while let Some(signal) = signals.next().await {
  59. match signal {
  60. SIGINT => {
  61. stop_s.send(()).await?;
  62. break
  63. }
  64. _ => unreachable!(),
  65. }
  66. }
  67. Ok::<(), crate::Error>(())
  68. });
  69. loop {
  70. let event = futures::select! {
  71. msg = self.socket.recv().fuse() => NetEvent::Receive(msg?),
  72. msg = self.recv_queue.recv().fuse() => NetEvent::Send(msg?),
  73. _ = stop_r.recv().fuse() => NetEvent::Stop
  74. };
  75. match event {
  76. NetEvent::Receive(msg) => {
  77. if let Some(peer) = msg.get(0) {
  78. if let Some(request) = msg.get(1) {
  79. let request: Vec<u8> = request.to_vec();
  80. let request: Request = deserialize(&request)?;
  81. self.send_queue.send((peer.to_vec(), request)).await?;
  82. }
  83. }
  84. }
  85. NetEvent::Send((peer, reply)) => {
  86. let peer = Bytes::from(peer);
  87. let mut msg: Vec<Bytes> = vec![peer];
  88. let reply: Vec<u8> = serialize(&reply);
  89. let reply = Bytes::from(reply);
  90. msg.push(reply);
  91. let reply = zeromq::ZmqMessage::try_from(msg)
  92. .map_err(|_| crate::Error::TryFromError)?;
  93. self.socket.send(reply).await?;
  94. }
  95. NetEvent::Stop => break,
  96. }
  97. }
  98. handle.close();
  99. signals_task.await?;
  100. debug!(target: "REP PROTOCOL API","{} SERVICE: Stopped", self.service_name);
  101. Ok(())
  102. }
  103. }
  104. pub struct ReqProtocol {
  105. addr: SocketAddr,
  106. socket: zeromq::DealerSocket,
  107. service_name: String,
  108. }
  109. impl ReqProtocol {
  110. pub fn new(addr: SocketAddr, service_name: String) -> ReqProtocol {
  111. let socket = zeromq::DealerSocket::new();
  112. ReqProtocol { addr, socket, service_name }
  113. }
  114. pub async fn start(&mut self) -> Result<()> {
  115. let addr = addr_to_string(self.addr);
  116. self.socket.connect(addr.as_str()).await?;
  117. debug!(target: "REQ PROTOCOL API","{} SERVICE: Connected To {}", self.service_name, self.addr);
  118. Ok(())
  119. }
  120. pub async fn request(
  121. &mut self,
  122. command: u8,
  123. data: Vec<u8>,
  124. handle_error: Arc<dyn Fn(u32) + Send + Sync>,
  125. ) -> Result<Option<Vec<u8>>> {
  126. let request = Request::new(command, data);
  127. let req = serialize(&request);
  128. let req = bytes::Bytes::from(req);
  129. let req: zeromq::ZmqMessage = req.into();
  130. self.socket.send(req).await?;
  131. debug!(
  132. target: "REQ PROTOCOL API",
  133. "{} SERVICE: Sent Request {{ command: {} }}",
  134. self.service_name, command
  135. );
  136. let rep: zeromq::ZmqMessage = self.socket.recv().await?;
  137. if let Some(reply) = rep.get(0) {
  138. let reply: Vec<u8> = reply.to_vec();
  139. let reply: Reply = deserialize(&reply)?;
  140. debug!(
  141. target: "REQ PROTOCOL API",
  142. "{} SERVICE: Received Reply {{ error: {} }}",
  143. self.service_name,
  144. reply.has_error()
  145. );
  146. if reply.has_error() {
  147. handle_error(reply.get_error());
  148. return Ok(None)
  149. }
  150. if reply.get_id() != request.get_id() {
  151. warn!("Reply id is not equal to Request id");
  152. return Ok(None)
  153. }
  154. Ok(Some(reply.get_payload()))
  155. } else {
  156. Err(crate::Error::ZmqError("Couldn't parse ZmqMessage".to_string()))
  157. }
  158. }
  159. }
  160. pub struct Publisher {
  161. addr: SocketAddr,
  162. socket: zeromq::PubSocket,
  163. service_name: String,
  164. }
  165. impl Publisher {
  166. pub fn new(addr: SocketAddr, service_name: String) -> Publisher {
  167. let socket = zeromq::PubSocket::new();
  168. Publisher { addr, socket, service_name }
  169. }
  170. pub async fn start(&mut self, recv_queue: async_channel::Receiver<Vec<u8>>) -> Result<()> {
  171. let addr = addr_to_string(self.addr);
  172. self.socket.bind(addr.as_str()).await?;
  173. debug!(
  174. target: "PUBLISHER API",
  175. "{} SERVICE : Bound To {}",
  176. self.service_name, addr
  177. );
  178. loop {
  179. let msg = recv_queue.recv().await?;
  180. self.publish(msg).await?;
  181. }
  182. }
  183. async fn publish(&mut self, data: Vec<u8>) -> Result<()> {
  184. let data = Bytes::from(data);
  185. self.socket.send(data.into()).await?;
  186. Ok(())
  187. }
  188. }
  189. pub struct Subscriber {
  190. addr: SocketAddr,
  191. socket: zeromq::SubSocket,
  192. service_name: String,
  193. }
  194. impl Subscriber {
  195. pub fn new(addr: SocketAddr, service_name: String) -> Subscriber {
  196. let socket = zeromq::SubSocket::new();
  197. Subscriber { addr, socket, service_name }
  198. }
  199. pub async fn start(&mut self) -> Result<()> {
  200. let addr = addr_to_string(self.addr);
  201. self.socket.connect(addr.as_str()).await?;
  202. self.socket.subscribe("").await?;
  203. debug!(
  204. target: "SUBSCRIBER API",
  205. "{} SERVICE : Connected To {}",
  206. self.service_name, addr
  207. );
  208. Ok(())
  209. }
  210. pub async fn fetch<T: Decodable>(&mut self) -> Result<T> {
  211. let data = self.socket.recv().await?;
  212. match data.get(0) {
  213. Some(d) => {
  214. let data = d.to_vec();
  215. let data: T = deserialize(&data)?;
  216. Ok(data)
  217. }
  218. None => Err(crate::Error::ZmqError("Couldn't parse ZmqMessage".to_string())),
  219. }
  220. }
  221. }
  222. #[derive(Debug, PartialEq)]
  223. pub struct Request {
  224. command: u8,
  225. id: u32,
  226. payload: Vec<u8>,
  227. }
  228. impl Request {
  229. pub fn new(command: u8, payload: Vec<u8>) -> Request {
  230. let id = Self::gen_id();
  231. Request { command, id, payload }
  232. }
  233. fn gen_id() -> u32 {
  234. let mut rng = rand::thread_rng();
  235. rng.gen()
  236. }
  237. pub fn get_id(&self) -> u32 {
  238. self.id
  239. }
  240. pub fn get_command(&self) -> u8 {
  241. self.command
  242. }
  243. pub fn get_payload(&self) -> Vec<u8> {
  244. self.payload.clone()
  245. }
  246. }
  247. #[derive(Debug, PartialEq)]
  248. pub struct Reply {
  249. id: u32,
  250. error: u32,
  251. payload: Vec<u8>,
  252. }
  253. impl Reply {
  254. pub fn from(request: &Request, error: u32, payload: Vec<u8>) -> Reply {
  255. Reply { id: request.get_id(), error, payload }
  256. }
  257. pub fn has_error(&self) -> bool {
  258. self.error != 0
  259. }
  260. pub fn get_error(&self) -> u32 {
  261. self.error
  262. }
  263. pub fn get_payload(&self) -> Vec<u8> {
  264. self.payload.clone()
  265. }
  266. pub fn set_payload(&mut self, payload: Vec<u8>) {
  267. self.payload = payload;
  268. }
  269. pub fn set_error(&mut self, error: u32) {
  270. self.error = error;
  271. }
  272. pub fn get_id(&self) -> u32 {
  273. self.id
  274. }
  275. }
  276. impl Encodable for Request {
  277. fn encode<S: io::Write>(&self, mut s: S) -> Result<usize> {
  278. let mut len = 0;
  279. len += self.command.encode(&mut s)?;
  280. len += self.id.encode(&mut s)?;
  281. len += self.payload.encode(&mut s)?;
  282. Ok(len)
  283. }
  284. }
  285. impl Encodable for Reply {
  286. fn encode<S: io::Write>(&self, mut s: S) -> Result<usize> {
  287. let mut len = 0;
  288. len += self.id.encode(&mut s)?;
  289. len += self.error.encode(&mut s)?;
  290. len += self.payload.encode(&mut s)?;
  291. Ok(len)
  292. }
  293. }
  294. impl Decodable for Request {
  295. fn decode<D: io::Read>(mut d: D) -> Result<Self> {
  296. Ok(Self {
  297. command: Decodable::decode(&mut d)?,
  298. id: Decodable::decode(&mut d)?,
  299. payload: Decodable::decode(&mut d)?,
  300. })
  301. }
  302. }
  303. impl Decodable for Reply {
  304. fn decode<D: io::Read>(mut d: D) -> Result<Self> {
  305. Ok(Self {
  306. id: Decodable::decode(&mut d)?,
  307. error: Decodable::decode(&mut d)?,
  308. payload: Decodable::decode(&mut d)?,
  309. })
  310. }
  311. }
  312. #[cfg(test)]
  313. mod tests {
  314. use super::{Reply, Request, Result};
  315. use crate::serial::{deserialize, serialize};
  316. #[test]
  317. fn serialize_and_deserialize_request_test() {
  318. let request = Request::new(2, vec![2, 3, 4, 6, 4]);
  319. let serialized_request = serialize(&request);
  320. assert!((deserialize(&serialized_request) as Result<bool>).is_err());
  321. let deserialized_request = deserialize(&serialized_request).ok();
  322. assert_eq!(deserialized_request, Some(request));
  323. }
  324. #[test]
  325. fn serialize_and_deserialize_reply_test() {
  326. let request = Request::new(2, vec![2, 3, 4, 6, 4]);
  327. let reply = Reply::from(&request, 0, vec![2, 3, 4, 6, 4]);
  328. let serialized_reply = serialize(&reply);
  329. assert!((deserialize(&serialized_reply) as Result<bool>).is_err());
  330. let deserialized_reply = deserialize(&serialized_reply).ok();
  331. assert_eq!(deserialized_reply, Some(reply));
  332. }
  333. }