Skip to main content

redevplugin_worker_sdk/
websocket.rs

1use crate::api;
2use crate::error::{Error, Result};
3use crate::http::Header;
4use crate::resource::{
5    Handle, IO_FLAG_BINARY, IO_FLAG_EOF, IO_FLAG_MESSAGE_END, IO_FLAG_TEXT, MAX_IO_CHUNK_BYTES,
6};
7use serde::{Deserialize, Serialize};
8
9#[derive(Debug, Clone, Serialize)]
10pub struct WebSocketOpen {
11    pub url: String,
12    #[serde(default)]
13    pub headers: Vec<Header>,
14    #[serde(default)]
15    pub subprotocols: Vec<String>,
16    #[serde(default)]
17    #[serde(skip_serializing_if = "Option::is_none")]
18    pub timeout_ms: Option<u32>,
19}
20
21pub type OpenOptions = WebSocketOpen;
22
23#[derive(Deserialize)]
24struct OpenResult {
25    handle: u64,
26    #[serde(default)]
27    protocol: String,
28    #[serde(default)]
29    response_headers: Vec<Header>,
30}
31
32#[derive(Debug, Clone, PartialEq, Eq)]
33pub enum Message {
34    Text(String),
35    Binary(Vec<u8>),
36}
37
38pub struct WebSocket {
39    handle: Handle,
40    pub protocol: String,
41    pub response_headers: Vec<Header>,
42}
43
44impl WebSocket {
45    pub fn open(options: WebSocketOpen) -> Result<Self> {
46        let opened: OpenResult = api::call("net.websocket.open", &options)?;
47        Ok(Self {
48            handle: Handle::new(opened.handle)?,
49            protocol: opened.protocol,
50            response_headers: opened.response_headers,
51        })
52    }
53
54    pub fn send_text(&mut self, text: &str) -> Result<()> {
55        self.send_message(text.as_bytes(), IO_FLAG_TEXT)
56    }
57
58    pub fn send_binary(&mut self, bytes: &[u8]) -> Result<()> {
59        self.send_message(bytes, IO_FLAG_BINARY)
60    }
61
62    fn send_message(&mut self, bytes: &[u8], kind: u32) -> Result<()> {
63        if bytes.is_empty() {
64            return self.handle.write(bytes, kind | IO_FLAG_MESSAGE_END);
65        }
66        let chunk_count = bytes.len().div_ceil(MAX_IO_CHUNK_BYTES);
67        for (index, chunk) in bytes.chunks(MAX_IO_CHUNK_BYTES).enumerate() {
68            let mut flags = if index == 0 { kind } else { 0 };
69            if index + 1 == chunk_count {
70                flags |= IO_FLAG_MESSAGE_END;
71            }
72            self.handle.write(chunk, flags)?;
73        }
74        Ok(())
75    }
76
77    pub fn receive(&mut self) -> Result<Message> {
78        let mut body = Vec::new();
79        let mut kind = 0;
80        loop {
81            let (chunk, flags) = self.handle.read(MAX_IO_CHUNK_BYTES)?;
82            let message_kind = flags & (IO_FLAG_TEXT | IO_FLAG_BINARY);
83            if flags & IO_FLAG_EOF != 0 {
84                return Err(Error::from_abi_status(-5));
85            }
86            if message_kind == IO_FLAG_TEXT | IO_FLAG_BINARY {
87                return Err(Error::internal("WebSocket message type is invalid"));
88            }
89            if message_kind != 0 {
90                if kind != 0 {
91                    return Err(Error::internal(
92                        "WebSocket message type changed mid-message",
93                    ));
94                }
95                kind = message_kind;
96            }
97            body.extend_from_slice(&chunk);
98            if flags & IO_FLAG_MESSAGE_END != 0 {
99                return match kind {
100                    IO_FLAG_TEXT => String::from_utf8(body)
101                        .map(Message::Text)
102                        .map_err(|_| Error::internal("WebSocket text message is not UTF-8")),
103                    IO_FLAG_BINARY => Ok(Message::Binary(body)),
104                    _ => Err(Error::internal("WebSocket message type is missing")),
105                };
106            }
107            if chunk.is_empty() {
108                return Err(Error::internal("WebSocket receive made no progress"));
109            }
110        }
111    }
112
113    pub fn ping(&mut self) -> Result<()> {
114        #[derive(Serialize)]
115        struct Arguments {
116            handle: u64,
117        }
118        let _: serde_json::Value = api::call(
119            "net.websocket.ping",
120            &Arguments {
121                handle: self.handle.id(),
122            },
123        )?;
124        Ok(())
125    }
126
127    pub fn close(mut self, code: u16, reason: &str) -> Result<()> {
128        #[derive(Serialize)]
129        struct Arguments<'a> {
130            handle: u64,
131            code: u16,
132            reason: &'a str,
133        }
134        let _: serde_json::Value = api::call(
135            "net.websocket.close",
136            &Arguments {
137                handle: self.handle.id(),
138                code,
139                reason,
140            },
141        )?;
142        self.handle.disarm();
143        Ok(())
144    }
145}