Skip to main content

volans_ws/
lib.rs

1mod stream;
2
3use std::{
4    pin::Pin,
5    task::{Context, Poll},
6};
7
8use async_tungstenite::{
9    accept_async_with_config, client_async_with_config,
10    tungstenite::{self, http::Uri, protocol::WebSocketConfig},
11};
12use futures::{FutureExt, TryFutureExt};
13use stream::RwStreamSink;
14use volans_core::{
15    Listener, ListenerEvent, Multiaddr, Transport, TransportError, multiaddr::Protocol,
16};
17use volans_tcp::TcpStream;
18
19use crate::framed::BytesWebSocketStream;
20pub use tungstenite::Error;
21
22mod framed;
23
24#[derive(Debug, Clone)]
25pub struct Config {
26    pub websocket: WebSocketConfig,
27    pub tcp: volans_tcp::Config,
28}
29
30impl Default for Config {
31    fn default() -> Self {
32        Self::new()
33    }
34}
35
36impl Config {
37    pub fn new() -> Self {
38        Self {
39            websocket: WebSocketConfig::default(),
40            tcp: volans_tcp::Config::default(),
41        }
42    }
43
44    /// Set [`Self::read_buffer_size`].
45    pub fn read_buffer_size(mut self, read_buffer_size: usize) -> Self {
46        self.websocket.read_buffer_size = read_buffer_size;
47        self
48    }
49
50    /// Set [`Self::write_buffer_size`].
51    pub fn write_buffer_size(mut self, write_buffer_size: usize) -> Self {
52        self.websocket.write_buffer_size = write_buffer_size;
53        self
54    }
55
56    /// Set [`Self::max_write_buffer_size`].
57    pub fn max_write_buffer_size(mut self, max_write_buffer_size: usize) -> Self {
58        self.websocket.max_write_buffer_size = max_write_buffer_size;
59        self
60    }
61
62    /// Set [`Self::max_message_size`].
63    pub fn max_message_size(mut self, max_message_size: Option<usize>) -> Self {
64        self.websocket.max_message_size = max_message_size;
65        self
66    }
67
68    /// Set [`Self::max_frame_size`].
69    pub fn max_frame_size(mut self, max_frame_size: Option<usize>) -> Self {
70        self.websocket.max_frame_size = max_frame_size;
71        self
72    }
73
74    /// Set [`Self::accept_unmasked_frames`].
75    pub fn accept_unmasked_frames(mut self, accept_unmasked_frames: bool) -> Self {
76        self.websocket.accept_unmasked_frames = accept_unmasked_frames;
77        self
78    }
79}
80
81type ListenerUpgrade = Pin<
82    Box<dyn Future<Output = Result<RwStreamSink<BytesWebSocketStream<TcpStream>>, Error>> + Send>,
83>;
84
85impl Transport for Config {
86    type Output = RwStreamSink<BytesWebSocketStream<TcpStream>>;
87    type Error = tungstenite::Error;
88    type Dial = Pin<Box<dyn Future<Output = Result<Self::Output, Self::Error>> + Send>>;
89    type Incoming = ListenerUpgrade;
90    type Listener = ListenStream;
91
92    fn dial(&self, addr: Multiaddr) -> Result<Self::Dial, TransportError<Self::Error>> {
93        let config = self.websocket.clone();
94        tracing::debug!("Connecting to WebSocket at {}", addr);
95        let ws_addr =
96            parse_ws_dial_addr(&addr).map_err(|_| TransportError::NotSupported(addr.clone()))?;
97
98        let request = Uri::builder()
99            .scheme(if ws_addr.use_tls { "wss" } else { "ws" })
100            .authority(ws_addr.host_port.as_str())
101            .path_and_query(ws_addr.path.as_str())
102            .build()
103            .map_err(|_| TransportError::NotSupported(addr.clone()))?;
104
105        tracing::debug!("Connecting to WebSocket at {}", request);
106
107        let dialer = self
108            .tcp
109            .dial(ws_addr.tcp_addr)
110            .map_err(|e| e.map(tungstenite::Error::from))?;
111
112        Ok(dialer
113            .map_err(tungstenite::Error::from)
114            .and_then(move |stream| client_async_with_config(request, stream, Some(config)))
115            .map_ok(|(s, response)| {
116                tracing::debug!("WebSocket handshake response: {:?}", response);
117                BytesWebSocketStream::new(s)
118            })
119            .map_ok(RwStreamSink::new)
120            .boxed())
121    }
122
123    fn listen(&self, addr: Multiaddr) -> Result<Self::Listener, TransportError<Self::Error>> {
124        let (inner_addr, path) = parse_ws_listen_addr(&addr)
125            .ok_or_else(|| TransportError::NotSupported(addr.clone()))?;
126        let listener = self
127            .tcp
128            .listen(inner_addr)
129            .map_err(|e| e.map(tungstenite::Error::from))?;
130        tracing::debug!("Listening for WebSocket connections on {}", addr);
131        Ok(ListenStream {
132            path: path.map(|r| r.to_string()),
133            config: self.websocket.clone(),
134            inner: listener,
135        })
136    }
137}
138
139#[pin_project::pin_project]
140pub struct ListenStream {
141    path: Option<String>,
142    config: WebSocketConfig,
143    #[pin]
144    inner: volans_tcp::ListenStream,
145}
146
147fn append_on_addr(mut addr: Multiaddr, path: Option<&str>) -> Multiaddr {
148    addr.push(Protocol::Ws);
149    if let Some(path) = path {
150        addr.push(Protocol::Path(path.into()));
151    }
152    addr
153}
154
155impl Listener for ListenStream {
156    type Output = RwStreamSink<BytesWebSocketStream<TcpStream>>;
157    type Error = tungstenite::Error;
158    type Upgrade = ListenerUpgrade;
159
160    fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
161        let this = self.project();
162        this.inner.poll_close(cx).map_err(tungstenite::Error::from)
163    }
164
165    fn poll_event(
166        self: Pin<&mut Self>,
167        cx: &mut Context<'_>,
168    ) -> Poll<ListenerEvent<Self::Upgrade, Self::Error>> {
169        let this = self.project();
170
171        let inner_event = {
172            match this.inner.poll_event(cx) {
173                Poll::Ready(event) => event,
174                Poll::Pending => return Poll::Pending,
175            }
176        };
177
178        let event = match inner_event {
179            ListenerEvent::AddressExpired(addr) => {
180                ListenerEvent::AddressExpired(append_on_addr(addr, this.path.as_deref()))
181            }
182            ListenerEvent::NewAddress(multiaddr) => {
183                ListenerEvent::NewAddress(append_on_addr(multiaddr, this.path.as_deref()))
184            }
185            ListenerEvent::Incoming {
186                local_addr,
187                remote_addr,
188                upgrade,
189            } => {
190                let config = this.config.clone();
191                let upgrade = upgrade
192                    .map_err(Error::from)
193                    .and_then(move |stream| {
194                        accept_async_with_config(stream, Some(config))
195                            .map_ok(BytesWebSocketStream::new)
196                            .map_ok(RwStreamSink::new)
197                    })
198                    .boxed();
199                ListenerEvent::Incoming {
200                    local_addr: append_on_addr(local_addr, this.path.as_deref()),
201                    remote_addr: append_on_addr(remote_addr, this.path.as_deref()),
202                    upgrade,
203                }
204            }
205            ListenerEvent::Closed(r) => ListenerEvent::Closed(r.map_err(Error::from)),
206            ListenerEvent::Error(err) => ListenerEvent::Error(err.into()),
207        };
208        Poll::Ready(event)
209    }
210}
211
212fn parse_ws_listen_addr(addr: &Multiaddr) -> Option<(Multiaddr, Option<String>)> {
213    let mut inner_addr = addr.clone();
214    let maybe_path = inner_addr.pop()?;
215    match maybe_path {
216        Protocol::Path(path) => match inner_addr.pop()? {
217            Protocol::Ws => Some((inner_addr, Some(path.to_string()))),
218            _ => None,
219        },
220        Protocol::Ws => Some((inner_addr, None)),
221        _ => None,
222    }
223}
224
225fn parse_ws_dial_addr(addr: &Multiaddr) -> Result<WsAddress, ()> {
226    let mut protocols = addr.iter();
227    let mut ip = protocols.next();
228    let mut tcp = protocols.next();
229
230    let (host_port, server_name) = loop {
231        match (ip, tcp) {
232            (Some(Protocol::Ip4(ip)), Some(Protocol::Tcp(port))) => {
233                let host_port = format!("{}:{}", ip, port);
234                break (host_port, ip.to_string());
235            }
236            (Some(Protocol::Ip6(ip)), Some(Protocol::Tcp(port))) => {
237                break (format!("[{ip}]:{port}"), ip.to_string());
238            }
239            (Some(Protocol::Dns(h)), Some(Protocol::Tcp(port)))
240            | (Some(Protocol::Dns4(h)), Some(Protocol::Tcp(port)))
241            | (Some(Protocol::Dns6(h)), Some(Protocol::Tcp(port))) => {
242                break (format!("{h}:{port}"), h.to_string());
243            }
244            (Some(_), Some(p)) => {
245                ip = Some(p);
246                tcp = protocols.next();
247            }
248            _ => return Err(()),
249        }
250    };
251
252    let mut protocols = addr.clone();
253    let mut peer = None;
254    let mut path = "/".to_string();
255    let (use_tls, path) = loop {
256        match protocols.pop() {
257            p @ Some(Protocol::Peer(_)) => peer = p,
258            Some(Protocol::Path(x_path)) => path = x_path.to_string(),
259            Some(Protocol::Ws) => match protocols.pop() {
260                Some(Protocol::Tls) => break (true, path),
261                Some(p) => {
262                    protocols.push(p);
263                    break (false, path);
264                }
265                None => return Err(()),
266            },
267            _ => return Err(()),
268        }
269    };
270    let tcp_addr = match peer {
271        Some(p) => protocols.with(p),
272        None => protocols,
273    };
274
275    Ok(WsAddress {
276        host_port,
277        server_name,
278        path,
279        use_tls,
280        tcp_addr,
281    })
282}
283
284#[derive(Debug)]
285struct WsAddress {
286    host_port: String,
287    server_name: String,
288    path: String,
289    use_tls: bool,
290    tcp_addr: Multiaddr,
291}