use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot, Mutex};
use tokio_tungstenite::connect_async;
use tokio_tungstenite::tungstenite::Message;
use futures_util::{StreamExt, SinkExt};
use uuid::Uuid;
#[derive(Debug, Deserialize, Serialize)]
struct RPCResult {
id: String,
result: Option<serde_json::Value>,
error: Option<serde_json::Value>,
}
#[derive(Debug, Deserialize, Serialize)]
struct RPCRequest {
id: String,
method: String,
params: serde_json::Value,
}
type RPCResultSender = oneshot::Sender<Result<RPCResult, String>>;
type RPCResultReceiver = oneshot::Receiver<Result<RPCResult, String>>;
struct RPCSession {
rpc_results: Arc<Mutex<HashMap<String, RPCResultSender>>>,
request_tx: mpsc::UnboundedSender<(String, serde_json::Value, RPCResultSender)>,
}
impl RPCSession {
async fn new(url: &str) -> Self {
let (ws_stream, _) = connect_async(url).await.expect("Failed to connect");
let (mut write, mut read) = ws_stream.split();
let rpc_results = Arc::new(Mutex::new(HashMap::new()));
let (request_tx, mut request_rx) = mpsc::unbounded_channel();
let rpc_results_clone: Arc<Mutex<HashMap<String, RPCResultSender>>> = Arc::clone(&rpc_results);
tokio::spawn(async move {
while let Some(msg) = read.next().await {
if let Ok(msg) = msg {
if let Ok(result) = serde_json::from_str::<RPCResult>(&msg.to_string()) {
if let Some(sender) = rpc_results_clone.lock().await.remove(&result.id) {
let _ = sender.send(Ok(result));
} else {
eprintln!("Unexpected ws msg: {}", msg);
}
} else {
eprintln!("Unhandled RPC msg!?\n{}", msg);
}
}
}
});
let rpc_results_clone: Arc<Mutex<HashMap<String, RPCResultSender>>> = Arc::clone(&rpc_results);
tokio::spawn(async move {
while let Some((method, params, sender)) = request_rx.recv().await {
let id = Uuid::new_v4().to_string();
let msg = RPCRequest {
id: id.clone(),
method,
params,
};
rpc_results_clone.lock().await.insert(id, sender);
let msg = Message::Text(serde_json::to_string(&msg).unwrap());
write.send(msg).await.unwrap();
}
});
Self {
rpc_results,
request_tx,
}
}
async fn send_request(&self, method: String, params: serde_json::Value) -> RPCResultReceiver {
let (sender, receiver) = oneshot::channel();
let _ = self.request_tx.send((method, params, sender));
receiver
}
}