messages.rs 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454
  1. use futures::prelude::*;
  2. use log::*;
  3. use num_enum::{IntoPrimitive, TryFromPrimitive};
  4. use smol::Executor;
  5. use smol::{Async, Timer};
  6. use std::convert::TryFrom;
  7. use std::io;
  8. use std::io::Cursor;
  9. use std::net::{SocketAddr, TcpStream};
  10. use std::sync::Arc;
  11. use std::time::Duration;
  12. use crate::async_serial::{AsyncReadExt, AsyncWriteExt};
  13. use crate::error::{Error, Result};
  14. pub use crate::net::AsyncTcpStream;
  15. use crate::serial::{serialize, Decodable, Encodable, VarInt};
  16. const MAGIC_BYTES: [u8; 4] = [0xd9, 0xef, 0xb6, 0x7d];
  17. pub type Ciphertext = Vec<u8>;
  18. pub type CiphertextHash = [u8; 32];
  19. // Packets and Message because Rust doesn't allow value
  20. // aliasing from ADL type enums (which Message uses).
  21. #[derive(IntoPrimitive, TryFromPrimitive, Copy, Clone, PartialEq, Eq, Hash)]
  22. #[repr(u8)]
  23. pub enum PacketType {
  24. Ping = 0,
  25. Pong = 1,
  26. GetAddrs = 2,
  27. Addrs = 3,
  28. Sync = 4,
  29. Inv = 5,
  30. GetSlabs = 6,
  31. Slab = 7,
  32. Version = 8,
  33. Verack = 9,
  34. }
  35. pub enum Message {
  36. Ping,
  37. Pong,
  38. GetAddrs(GetAddrsMessage),
  39. Addrs(AddrsMessage),
  40. Sync,
  41. Inv(InvMessage),
  42. GetSlabs(GetSlabsMessage),
  43. Slab(SlabMessage),
  44. Version(VersionMessage),
  45. Verack(VerackMessage),
  46. }
  47. pub struct GetAddrsMessage {}
  48. pub struct GetSlabsMessage {
  49. pub slabs_hash: Vec<[u8; 32]>,
  50. }
  51. #[derive(Clone)]
  52. pub struct SlabMessage {
  53. pub nonce: [u8; 12],
  54. pub ciphertext: Ciphertext,
  55. }
  56. pub struct InvMessage {
  57. pub slabs_hash: Vec<[u8; 32]>,
  58. }
  59. pub struct AddrsMessage {
  60. pub addrs: Vec<SocketAddr>,
  61. }
  62. pub struct VersionMessage {}
  63. pub struct VerackMessage {}
  64. impl Encodable for GetSlabsMessage {
  65. fn encode<S: io::Write>(&self, mut s: S) -> Result<usize> {
  66. let mut len = 0;
  67. len += self.slabs_hash.encode(&mut s)?;
  68. Ok(len)
  69. }
  70. }
  71. impl Decodable for GetSlabsMessage {
  72. fn decode<D: io::Read>(mut d: D) -> Result<Self> {
  73. Ok(Self {
  74. slabs_hash: Decodable::decode(&mut d)?,
  75. })
  76. }
  77. }
  78. impl Encodable for SlabMessage {
  79. fn encode<S: io::Write>(&self, mut s: S) -> Result<usize> {
  80. let mut len = 0;
  81. len += self.nonce.encode(&mut s)?;
  82. len += self.ciphertext.encode(&mut s)?;
  83. Ok(len)
  84. }
  85. }
  86. impl Decodable for SlabMessage {
  87. fn decode<D: io::Read>(mut d: D) -> Result<Self> {
  88. Ok(Self {
  89. nonce: Decodable::decode(&mut d)?,
  90. ciphertext: Decodable::decode(&mut d)?,
  91. })
  92. }
  93. }
  94. impl Encodable for InvMessage {
  95. fn encode<S: io::Write>(&self, mut s: S) -> Result<usize> {
  96. let mut len = 0;
  97. len += self.slabs_hash.encode(&mut s)?;
  98. Ok(len)
  99. }
  100. }
  101. impl Decodable for InvMessage {
  102. fn decode<D: io::Read>(mut d: D) -> Result<Self> {
  103. Ok(Self {
  104. slabs_hash: Decodable::decode(&mut d)?,
  105. })
  106. }
  107. }
  108. impl Encodable for GetAddrsMessage {
  109. fn encode<S: io::Write>(&self, mut _s: S) -> Result<usize> {
  110. let len = 0;
  111. Ok(len)
  112. }
  113. }
  114. impl Decodable for GetAddrsMessage {
  115. fn decode<D: io::Read>(mut _d: D) -> Result<Self> {
  116. Ok(Self {})
  117. }
  118. }
  119. impl Encodable for AddrsMessage {
  120. fn encode<S: io::Write>(&self, mut s: S) -> Result<usize> {
  121. let mut len = 0;
  122. len += self.addrs.encode(&mut s)?;
  123. Ok(len)
  124. }
  125. }
  126. impl Decodable for AddrsMessage {
  127. fn decode<D: io::Read>(mut d: D) -> Result<Self> {
  128. Ok(Self {
  129. addrs: Decodable::decode(&mut d)?,
  130. })
  131. }
  132. }
  133. impl Encodable for VersionMessage {
  134. fn encode<S: io::Write>(&self, mut s: S) -> Result<usize> {
  135. Ok(0)
  136. }
  137. }
  138. impl Decodable for VersionMessage {
  139. fn decode<D: io::Read>(mut d: D) -> Result<Self> {
  140. Ok(Self {})
  141. }
  142. }
  143. impl Encodable for VerackMessage {
  144. fn encode<S: io::Write>(&self, mut s: S) -> Result<usize> {
  145. Ok(0)
  146. }
  147. }
  148. impl Decodable for VerackMessage {
  149. fn decode<D: io::Read>(mut d: D) -> Result<Self> {
  150. Ok(Self {})
  151. }
  152. }
  153. impl Message {
  154. pub fn packet_type(&self) -> PacketType {
  155. match self {
  156. Message::Ping =>
  157. PacketType::Ping,
  158. Message::Pong =>
  159. PacketType::Pong,
  160. Message::GetAddrs(message) =>
  161. PacketType::GetAddrs,
  162. Message::Addrs(message) =>
  163. PacketType::Addrs,
  164. Message::Sync =>
  165. PacketType::Sync,
  166. Message::Inv(message) =>
  167. PacketType::Inv,
  168. Message::GetSlabs(message) =>
  169. PacketType::GetSlabs,
  170. Message::Slab(message) =>
  171. PacketType::Slab,
  172. Message::Version(message) =>
  173. PacketType::Version,
  174. Message::Verack(message) =>
  175. PacketType::Verack,
  176. }
  177. }
  178. pub fn pack(&self) -> Result<Packet> {
  179. match self {
  180. Message::Ping => Ok(Packet {
  181. command: PacketType::Ping,
  182. payload: Vec::new(),
  183. }),
  184. Message::Pong => Ok(Packet {
  185. command: PacketType::Pong,
  186. payload: Vec::new(),
  187. }),
  188. Message::GetAddrs(message) => {
  189. let mut payload = Vec::new();
  190. message.encode(&mut payload)?;
  191. Ok(Packet {
  192. command: PacketType::GetAddrs,
  193. payload,
  194. })
  195. }
  196. Message::Addrs(message) => {
  197. let mut payload = Vec::new();
  198. message.encode(Cursor::new(&mut payload))?;
  199. Ok(Packet {
  200. command: PacketType::Addrs,
  201. payload,
  202. })
  203. }
  204. Message::Sync => {
  205. let payload = Vec::new();
  206. Ok(Packet {
  207. command: PacketType::Sync,
  208. payload,
  209. })
  210. }
  211. Message::Inv(message) => {
  212. let payload = serialize(message);
  213. Ok(Packet {
  214. command: PacketType::Inv,
  215. payload,
  216. })
  217. }
  218. Message::GetSlabs(message) => {
  219. let payload = serialize(message);
  220. Ok(Packet {
  221. command: PacketType::GetSlabs,
  222. payload,
  223. })
  224. }
  225. Message::Slab(message) => {
  226. let payload = serialize(message);
  227. Ok(Packet {
  228. command: PacketType::Slab,
  229. payload,
  230. })
  231. }
  232. Message::Version(message) => {
  233. let payload = serialize(message);
  234. Ok(Packet {
  235. command: PacketType::Version,
  236. payload,
  237. })
  238. }
  239. Message::Verack(message) => {
  240. let payload = serialize(message);
  241. Ok(Packet {
  242. command: PacketType::Verack,
  243. payload,
  244. })
  245. }
  246. }
  247. }
  248. pub fn unpack(packet: Packet) -> Result<Self> {
  249. let cursor = Cursor::new(packet.payload.clone());
  250. match packet.command {
  251. PacketType::Ping => Ok(Self::Ping),
  252. PacketType::Pong => Ok(Self::Pong),
  253. PacketType::GetAddrs => Ok(Self::GetAddrs(GetAddrsMessage::decode(cursor)?)),
  254. PacketType::Addrs => Ok(Self::Addrs(AddrsMessage::decode(cursor)?)),
  255. PacketType::Sync => Ok(Self::Sync),
  256. PacketType::Inv => Ok(Self::Inv(InvMessage::decode(cursor)?)),
  257. PacketType::GetSlabs => Ok(Self::GetSlabs(GetSlabsMessage::decode(cursor)?)),
  258. PacketType::Slab => Ok(Self::Slab(SlabMessage::decode(cursor)?)),
  259. PacketType::Version => Ok(Self::Version(VersionMessage::decode(cursor)?)),
  260. PacketType::Verack => Ok(Self::Verack(VerackMessage::decode(cursor)?)),
  261. }
  262. }
  263. pub fn name(&self) -> &'static str {
  264. match self {
  265. Message::Ping => "Ping",
  266. Message::Pong => "Pong",
  267. Message::GetAddrs(_) => "GetAddrs",
  268. Message::Addrs(_) => "Addrs",
  269. Message::Sync => "Sync",
  270. Message::Inv(_) => "Inv",
  271. Message::GetSlabs(_) => "GetSlabs",
  272. Message::Slab(_) => "Slab",
  273. Message::Version(_) => "Version",
  274. Message::Verack(_) => "Verack",
  275. }
  276. }
  277. }
  278. // Packets are the base type read from the network
  279. // These are converted to messages and passed to event loop
  280. pub struct Packet {
  281. pub command: PacketType,
  282. pub payload: Vec<u8>,
  283. }
  284. pub async fn read_packet<R: AsyncRead + Unpin>(stream: &mut R) -> Result<Packet> {
  285. // Packets have a 4 byte header of magic digits
  286. // This is used for network debugging
  287. let mut magic = [0u8; 4];
  288. stream.read_exact(&mut magic).await?;
  289. //debug!("read magic {:?}", magic);
  290. if magic != MAGIC_BYTES {
  291. return Err(Error::MalformedPacket);
  292. }
  293. // The type of the message
  294. let command = AsyncReadExt::read_u8(stream).await?;
  295. //debug!("read command: {}", command);
  296. let command = PacketType::try_from(command).map_err(|_| Error::MalformedPacket)?;
  297. let payload_len = VarInt::decode_async(stream).await?.0 as usize;
  298. // The message-dependent data (see message types)
  299. let mut payload = vec![0u8; payload_len];
  300. stream.read_exact(&mut payload).await?;
  301. //debug!("read payload");
  302. Ok(Packet { command, payload })
  303. }
  304. pub async fn send_packet<W: AsyncWrite + Unpin>(stream: &mut W, packet: Packet) -> Result<()> {
  305. stream.write_all(&MAGIC_BYTES).await?;
  306. AsyncWriteExt::write_u8(stream, packet.command as u8).await?;
  307. assert_eq!(std::mem::size_of::<usize>(), std::mem::size_of::<u64>());
  308. VarInt(packet.payload.len() as u64)
  309. .encode_async(stream)
  310. .await?;
  311. stream.write_all(&packet.payload).await?;
  312. Ok(())
  313. }
  314. pub async fn receive_message<R: AsyncRead + Unpin>(stream: &mut R) -> Result<Message> {
  315. let packet = read_packet(stream).await?;
  316. let message = Message::unpack(packet)?;
  317. debug!("received Message::{}", message.name());
  318. Ok(message)
  319. }
  320. pub async fn send_message<W: AsyncWrite + Unpin>(stream: &mut W, message: Message) -> Result<()> {
  321. debug!("sending Message::{}", message.name());
  322. let packet = message.pack()?;
  323. send_packet(stream, packet).await
  324. }
  325. // Eventloop event
  326. pub enum Event {
  327. // Message to be sent from event queue
  328. Send(Message),
  329. // Received message to process by protocol
  330. Receive(Message),
  331. // Connection ping-pong timeout
  332. Timeout,
  333. }
  334. pub async fn select_event(
  335. stream: &mut AsyncTcpStream,
  336. send_rx: &async_channel::Receiver<Message>,
  337. inactivity_timer: &InactivityTimer,
  338. ) -> Result<Event> {
  339. Ok(futures::select! {
  340. message = send_rx.recv().fuse() => Event::Send(message?),
  341. message = receive_message(stream).fuse() => Event::Receive(message?),
  342. _ = inactivity_timer.wait_for_wakeup().fuse() => Event::Timeout
  343. })
  344. }
  345. pub async fn sleep(seconds: u64) {
  346. Timer::after(Duration::from_secs(seconds)).await;
  347. }
  348. // Used for ping pong loop timer
  349. pub struct InactivityTimer {
  350. reset_sender: async_channel::Sender<()>,
  351. timeout_receiver: async_channel::Receiver<()>,
  352. task: smol::Task<()>,
  353. }
  354. impl InactivityTimer {
  355. pub fn new(executor: Arc<Executor<'_>>) -> Self {
  356. let (reset_sender, reset_receiver) = async_channel::bounded::<()>(1);
  357. let (timeout_sender, timeout_receiver) = async_channel::bounded::<()>(1);
  358. let task = executor.spawn(async {
  359. match Self::_start(reset_receiver, timeout_sender).await {
  360. Ok(()) => {}
  361. Err(err) => error!("InactivityTimer fatal error {}", err),
  362. }
  363. });
  364. Self {
  365. reset_sender,
  366. timeout_receiver,
  367. task,
  368. }
  369. }
  370. pub async fn stop(self) {
  371. self.task.cancel().await;
  372. }
  373. // This loop basically waits for 10 secs. If it doesn't
  374. // receive a signal that something happened then it will
  375. // send a timeout signal. This will wakeup the main event loop
  376. // and the connection will be dropped.
  377. async fn _start(
  378. reset_rx: async_channel::Receiver<()>,
  379. timeout_sx: async_channel::Sender<()>,
  380. ) -> Result<()> {
  381. loop {
  382. let is_awake = futures::select! {
  383. _ = reset_rx.recv().fuse() => true,
  384. _ = sleep(10).fuse() => false
  385. };
  386. if !is_awake {
  387. warn!("InactivityTimer timeout");
  388. timeout_sx.send(()).await?;
  389. }
  390. }
  391. }
  392. pub async fn reset(&self) -> Result<()> {
  393. self.reset_sender.send(()).await?;
  394. Ok(())
  395. }
  396. pub async fn wait_for_wakeup(&self) -> Result<()> {
  397. Ok(self.timeout_receiver.recv().await?)
  398. }
  399. }