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 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 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 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 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 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 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}