volt-client-rs 0.1.12

Volt websocket client library
Documentation
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> {
        // Encode the payload using the wasm protocol manager.
        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());

        // Send the request via the websocket send channel.
        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(),
                        ));
                    }
                }
            }
        }
    }
}