rpcserver.rs 5.0 KB

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