channel.rs 7.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225
  1. use async_std::sync::Mutex;
  2. use futures::io::{ReadHalf, WriteHalf};
  3. use futures::AsyncReadExt;
  4. use log::*;
  5. use smol::{Async, Executor};
  6. use std::net::{SocketAddr, TcpStream};
  7. use std::sync::atomic::{AtomicBool, Ordering};
  8. use std::sync::Arc;
  9. use crate::error;
  10. use crate::net::error::{NetError, NetResult};
  11. use crate::net::message_subscriber::{MessageSubscription, MessageSubsystem};
  12. use crate::net::messages;
  13. use crate::system::{StoppableTask, StoppableTaskPtr, Subscriber, SubscriberPtr, Subscription};
  14. pub type ChannelPtr = Arc<Channel>;
  15. pub struct Channel {
  16. reader: Mutex<ReadHalf<Async<TcpStream>>>,
  17. writer: Mutex<WriteHalf<Async<TcpStream>>>,
  18. address: SocketAddr,
  19. message_subsystem: MessageSubsystem,
  20. stop_subscriber: SubscriberPtr<NetError>,
  21. receive_task: StoppableTaskPtr,
  22. stopped: AtomicBool,
  23. }
  24. impl Channel {
  25. pub async fn new(
  26. stream: Async<TcpStream>,
  27. address: SocketAddr,
  28. ) -> Arc<Self> {
  29. let (reader, writer) = stream.split();
  30. let reader = Mutex::new(reader);
  31. let writer = Mutex::new(writer);
  32. let message_subsystem = MessageSubsystem::new();
  33. Self::setup_dispatchers(&message_subsystem).await;
  34. Arc::new(Self {
  35. reader,
  36. writer,
  37. address,
  38. message_subsystem,
  39. stop_subscriber: Subscriber::new(),
  40. receive_task: StoppableTask::new(),
  41. stopped: AtomicBool::new(false),
  42. })
  43. }
  44. pub fn start(self: Arc<Self>, executor: Arc<Executor<'_>>) {
  45. debug!(target: "net", "Channel::start() [START, address={}]", self.address());
  46. let self2 = self.clone();
  47. self.receive_task.clone().start(
  48. self.clone().main_receive_loop(),
  49. // Ignore stop handler
  50. |result| self2.handle_stop(result),
  51. NetError::ServiceStopped,
  52. executor,
  53. );
  54. debug!(target: "net", "Channel::start() [END, address={}]", self.address());
  55. }
  56. pub async fn stop(&self) {
  57. debug!(target: "net", "Channel::stop() [START, address={}]", self.address());
  58. assert_eq!(self.stopped.load(Ordering::Relaxed), false);
  59. self.stopped.store(false, Ordering::Relaxed);
  60. self.stop_subscriber.notify(NetError::ChannelStopped).await;
  61. self.receive_task.stop().await;
  62. self.message_subsystem.trigger_error(NetError::ChannelStopped).await;
  63. debug!(target: "net", "Channel::stop() [END, address={}]", self.address());
  64. }
  65. pub async fn subscribe_stop(&self) -> Subscription<NetError> {
  66. debug!(target: "net",
  67. "Channel::subscribe_stop() [START, address={}]",
  68. self.address()
  69. );
  70. // TODO: this should check the stopped status
  71. // Call to receive should return ChannelStopped on newly created sub
  72. let sub = self.stop_subscriber.clone().subscribe().await;
  73. debug!(target: "net",
  74. "Channel::subscribe_stop() [END, address={}]",
  75. self.address()
  76. );
  77. sub
  78. }
  79. pub async fn send<M: messages::Message>(&self, message: M) -> NetResult<()> {
  80. debug!(target: "net",
  81. "Channel::send() [START, command={:?}, address={}]",
  82. M::name(),
  83. self.address()
  84. );
  85. if self.stopped.load(Ordering::Relaxed) {
  86. return Err(NetError::ChannelStopped);
  87. }
  88. // Catch failure and stop channel, return a net error
  89. let result = match self.send_message(message).await {
  90. Ok(()) => Ok(()),
  91. Err(err) => {
  92. error!("Channel send error for [{}]: {}", self.address(), err);
  93. self.stop().await;
  94. Err(NetError::ChannelStopped)
  95. }
  96. };
  97. debug!(target: "net",
  98. "Channel::send() [END, command={:?}, address={}]",
  99. M::name(),
  100. self.address()
  101. );
  102. result
  103. }
  104. async fn send_message<M: messages::Message>(&self, message: M) -> error::Result<()> {
  105. let mut payload = Vec::new();
  106. message.encode(&mut payload)?;
  107. let packet = messages::Packet {
  108. command: String::from(M::name()),
  109. payload,
  110. };
  111. let stream = &mut *self.writer.lock().await;
  112. messages::send_packet(stream, packet).await
  113. }
  114. pub async fn subscribe_msg<M: messages::Message>(&self) -> NetResult<MessageSubscription<M>> {
  115. debug!(target: "net",
  116. "Channel::subscribe_msg() [START, command={:?}, address={}]",
  117. M::name(),
  118. self.address()
  119. );
  120. let sub = self.message_subsystem.subscribe::<M>().await;
  121. debug!(target: "net",
  122. "Channel::subscribe_msg() [END, command={:?}, address={}]",
  123. M::name(),
  124. self.address()
  125. );
  126. sub
  127. }
  128. pub fn address(&self) -> SocketAddr {
  129. self.address
  130. }
  131. fn is_eof_error(err: &error::Error) -> bool {
  132. match err {
  133. error::Error::Io(io_err) => io_err.kind() == std::io::ErrorKind::UnexpectedEof,
  134. _ => false,
  135. }
  136. }
  137. async fn setup_dispatchers(message_subsystem: &MessageSubsystem) {
  138. message_subsystem
  139. .add_dispatch::<messages::VersionMessage>()
  140. .await;
  141. message_subsystem
  142. .add_dispatch::<messages::VerackMessage>()
  143. .await;
  144. message_subsystem
  145. .add_dispatch::<messages::PingMessage>()
  146. .await;
  147. message_subsystem
  148. .add_dispatch::<messages::PongMessage>()
  149. .await;
  150. message_subsystem
  151. .add_dispatch::<messages::GetAddrsMessage>()
  152. .await;
  153. message_subsystem
  154. .add_dispatch::<messages::AddrsMessage>()
  155. .await;
  156. }
  157. pub fn get_message_subsystem(&self) -> &MessageSubsystem {
  158. &self.message_subsystem
  159. }
  160. async fn main_receive_loop(self: Arc<Self>) -> NetResult<()> {
  161. debug!(target: "net",
  162. "Channel::receive_loop() [START, address={}]",
  163. self.address()
  164. );
  165. let reader = &mut *self.reader.lock().await;
  166. loop {
  167. let packet = match messages::read_packet(reader).await {
  168. Ok(packet) => packet,
  169. Err(err) => {
  170. if Self::is_eof_error(&err) {
  171. info!("Channel {} disconnected", self.address());
  172. } else {
  173. error!("Read error on channel: {}", err);
  174. }
  175. debug!(target: "net",
  176. "Channel::receive_loop() stopping channel {}",
  177. self.address()
  178. );
  179. self.stop().await;
  180. return Err(NetError::ChannelStopped);
  181. }
  182. };
  183. // Send result to our subscribers
  184. self.message_subsystem
  185. .notify(&packet.command, packet.payload)
  186. .await;
  187. }
  188. }
  189. async fn handle_stop(self: Arc<Self>, result: NetResult<()>) {
  190. debug!(target: "net", "Channel::handle_stop() [START, address={}]", self.address());
  191. match result {
  192. Ok(()) => panic!("Channel task should never complete without error status"),
  193. Err(err) => {
  194. // Send this error to all channel subscribers
  195. self.message_subsystem.trigger_error(err).await;
  196. }
  197. }
  198. debug!(target: "net", "Channel::handle_stop() [END, address={}]", self.address());
  199. }
  200. }