messages.rs 13 KB

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