use std::collections::HashMap;
use std::sync::{Arc, Mutex as StdMutex};
use tokio::sync::{mpsc, Mutex};
#[derive(Debug, Clone, Hash, Eq, PartialEq)]
pub enum EventType {
Data,
Error,
End,
}
#[derive(Debug, Clone)]
pub enum Event {
Data(serde_json::Value),
Error(String),
End,
}
type Callback = Arc<StdMutex<dyn Fn(Event) + Send + Sync>>;
#[derive(Clone)]
pub struct WebsocketRpc {
pub id: u64,
protocol_manager: Arc<Mutex<volt_ws_protocol::rpc_manager::RpcManager>>,
send_channel: mpsc::Sender<Vec<u8>>,
callbacks: Arc<StdMutex<HashMap<EventType, Vec<Callback>>>>,
}
impl WebsocketRpc {
pub fn new(
id: u64,
protocol_manager: Arc<Mutex<volt_ws_protocol::rpc_manager::RpcManager>>,
send_channel: mpsc::Sender<Vec<u8>>,
) -> Self {
WebsocketRpc {
id,
protocol_manager,
send_channel,
callbacks: Arc::new(StdMutex::new(HashMap::new())),
}
}
pub fn on<F>(&self, event_type: EventType, callback: F)
where
F: Fn(Event) + Send + Sync + 'static,
{
let mut callbacks = self.callbacks.lock().unwrap();
let entry = callbacks.entry(event_type).or_insert_with(Vec::new);
entry.push(Arc::new(StdMutex::new(callback)));
}
fn emit(&self, event: Event) {
let event_type = match &event {
Event::Data(_) => EventType::Data,
Event::Error(_) => EventType::Error,
Event::End => EventType::End,
};
if let Some(callbacks) = self.callbacks.lock().unwrap().get(&event_type) {
for callback in callbacks.iter() {
let callback = callback.clone();
let event = event.clone();
(callback.lock().unwrap())(event);
}
}
}
pub fn abort(&mut self, error: String) {
self.emit(Event::Error(error));
}
async fn send_internal(&self, payload: &str) -> Result<(), String> {
let protocol = self.protocol_manager.lock().await;
let encoded = match protocol.encode_payload(&self.id, payload) {
Ok(encoded) => encoded,
Err(e) => {
return Err(format!("Failed to encode payload: {}", e));
}
};
println!("Sending payload size: {}", encoded.len());
let send_result = self.send_channel.send(encoded).await;
match send_result {
Ok(_) => Ok(()),
Err(e) => {
return Err(format!("Failed to send payload: {}", e));
}
}
}
pub async fn send(&self, payload: &serde_json::Value) -> Result<(), String> {
let payload_json = match serde_json::to_string(&payload) {
Ok(payload_json) => payload_json,
Err(e) => {
return Err(format!("Failed to serialize payload: {}", e));
}
};
self.send_internal(&payload_json).await
}
pub async fn end(&self) -> Result<(), String> {
self.send_internal("").await
}
pub fn handle_response(&mut self, response: &serde_json::Value) {
if !response["error"].is_null() {
self.abort(response["error"].as_str().unwrap().to_string());
} else {
let response_payload = &response["payload"];
if response_payload.is_null() {
println!(
"Received response with no payload for method_id: {}",
self.id
);
} else {
let method_payload = &response_payload["method_payload"];
if method_payload.is_null() {
let method_end = &response_payload["method_end"];
if method_end.is_null() {
println!(
"Received response with no method_payload for method_id: {}",
self.id
);
} else {
println!("received method_end for {} {}", self.id, method_end);
if response["error"].is_null() {
self.emit(Event::End);
} else {
self.emit(Event::Error(
response["error"].as_str().unwrap().to_string(),
));
}
}
} else {
println!("received payload for {} {}", self.id, method_payload);
let payload_json = match method_payload["json_payload"].as_str() {
Some(json_payload) => json_payload,
None => {
self.abort("Received response with no json_payload".to_string());
return;
}
};
let payload: serde_json::Value = match serde_json::from_str(payload_json) {
Ok(payload) => payload,
Err(e) => {
self.abort(format!("Failed to parse json_payload: {}", e));
return;
}
};
if payload["status"].is_null() {
self.emit(Event::Data(payload));
} else {
self.emit(Event::Error(
payload["status"]["message"].as_str().unwrap().to_string(),
));
}
}
}
}
}
}