rpcserver.rs 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164
  1. use std::{
  2. net::{SocketAddr, TcpListener, TcpStream},
  3. path::PathBuf,
  4. sync::Arc,
  5. };
  6. use async_executor::Executor;
  7. use async_native_tls::{Identity, TlsAcceptor};
  8. use async_trait::async_trait;
  9. use log::{debug, error, info};
  10. use smol::{
  11. io::{AsyncReadExt, AsyncWriteExt},
  12. Async,
  13. };
  14. use crate::{
  15. util::rpc::jsonrpc::{JsonRequest, JsonResult},
  16. Result,
  17. };
  18. pub struct RpcServerConfig {
  19. pub socket_addr: SocketAddr,
  20. pub use_tls: bool,
  21. pub identity_path: PathBuf,
  22. pub identity_pass: String,
  23. }
  24. #[async_trait]
  25. pub trait RequestHandler: Sync + Send {
  26. async fn handle_request(&self, req: JsonRequest, executor: Arc<Executor<'_>>) -> JsonResult;
  27. }
  28. async fn serve(
  29. mut stream: Async<TcpStream>,
  30. tls: Option<TlsAcceptor>,
  31. rh: Arc<impl RequestHandler + 'static>,
  32. executor: Arc<Executor<'_>>,
  33. ) -> Result<()> {
  34. debug!(target: "RPC SERVER", "Accepted connection");
  35. let mut buf = [0; 2048];
  36. match tls {
  37. None => loop {
  38. let n = match stream.read(&mut buf).await {
  39. Ok(n) if n == 0 => {
  40. debug!(target: "RPC SERVER", "Closed connection");
  41. return Ok(())
  42. }
  43. Ok(n) => n,
  44. Err(e) => {
  45. debug!(target: "RPC SERVER", "Failed to read from socket: {:#?}", e);
  46. debug!(target: "RPC SERVER", "Closed connection");
  47. return Ok(())
  48. }
  49. };
  50. let r: JsonRequest = match serde_json::from_slice(&buf[0..n]) {
  51. Ok(r) => r,
  52. Err(e) => {
  53. debug!(target: "RPC SERVER", "Received invalid JSON: {:#?}", e);
  54. debug!(target: "RPC SERVER", "Closed connection");
  55. return Ok(())
  56. }
  57. };
  58. let reply = rh.handle_request(r, executor.clone()).await;
  59. let j = serde_json::to_string(&reply)?;
  60. debug!(target: "RPC", "<-- {}", j);
  61. if let Err(e) = stream.write_all(j.as_bytes()).await {
  62. debug!(target: "RPC SERVER", "Failed to write to socket: {:#?}", e);
  63. debug!(target: "RPC SERVER", "Closed connection");
  64. return Ok(())
  65. }
  66. },
  67. Some(tls) => match tls.accept(stream).await {
  68. Ok(mut stream) => loop {
  69. let n = match stream.read(&mut buf).await {
  70. Ok(n) if n == 0 => {
  71. debug!(target: "RPC SERVER", "Closed connection");
  72. return Ok(())
  73. }
  74. Ok(n) => n,
  75. Err(e) => {
  76. debug!(target: "RPC SERVER", "Failed to read from socket: {:#?}", e);
  77. debug!(target: "RPC SERVER", "Closed connection");
  78. return Ok(())
  79. }
  80. };
  81. let r: JsonRequest = match serde_json::from_slice(&buf[0..n]) {
  82. Ok(r) => r,
  83. Err(e) => {
  84. debug!(target: "RPC SERVER", "Received invalid JSON: {:#?}", e);
  85. debug!(target: "RPC SERVER", "Closed connection");
  86. return Ok(())
  87. }
  88. };
  89. let reply = rh.handle_request(r, executor.clone()).await;
  90. let j = serde_json::to_string(&reply)?;
  91. debug!(target: "RPC", "<-- {}", j);
  92. if let Err(e) = stream.write_all(j.as_bytes()).await {
  93. debug!(target: "RPC SERVER", "Failed to write to socket: {:#?}", e);
  94. return Ok(())
  95. }
  96. },
  97. Err(e) => {
  98. debug!(target: "RPC SERVER", "Failed to establish TLS connection: {:#}", e);
  99. Ok(())
  100. }
  101. },
  102. }
  103. }
  104. async fn listen(
  105. listener: Async<TcpListener>,
  106. tls: Option<TlsAcceptor>,
  107. rh: Arc<impl RequestHandler + 'static>,
  108. executor: Arc<Executor<'_>>,
  109. ) -> Result<()> {
  110. match &tls {
  111. None => {
  112. info!(target: "RPC SERVER", "Listening on tcp://{}", listener.get_ref().local_addr()?)
  113. }
  114. Some(_) => {
  115. info!(target: "RPC SERVER", "Listening on tls://{}", listener.get_ref().local_addr()?)
  116. }
  117. }
  118. let ex = executor.clone();
  119. loop {
  120. let (stream, _) = listener.accept().await?;
  121. let tls = tls.clone();
  122. let rh_c = rh.clone();
  123. let ex2 = ex.clone();
  124. ex.spawn(async move {
  125. if let Err(err) = serve(stream, tls, rh_c, ex2.clone()).await {
  126. error!(target: "RPC SERVER", "Connection error: {:#?}", err);
  127. }
  128. })
  129. .detach();
  130. }
  131. }
  132. pub async fn listen_and_serve(
  133. cfg: RpcServerConfig,
  134. rh: Arc<impl RequestHandler + 'static>,
  135. executor: Arc<Executor<'_>>,
  136. ) -> Result<()> {
  137. let tls: Option<TlsAcceptor> = if cfg.use_tls {
  138. let ident_bytes = std::fs::read(cfg.identity_path)?;
  139. let identity = Identity::from_pkcs12(&ident_bytes, &cfg.identity_pass)?;
  140. Some(TlsAcceptor::from(native_tls::TlsAcceptor::new(identity)?))
  141. } else {
  142. None
  143. };
  144. let listener = listen(Async::<TcpListener>::bind(cfg.socket_addr)?, tls, rh, executor);
  145. listener.await
  146. }