1use std::rc::Rc;
2
3use bytes::Bytes;
4use futures::lock::Mutex;
5use js_sys::Uint8Array;
6use url::Url;
7use wasm_bindgen::JsCast;
8use wasm_bindgen_futures::JsFuture;
9use web_sys::{
10 WebTransport, WebTransportBidirectionalStream, WebTransportCloseInfo,
11 WebTransportReceiveStream, WebTransportSendStream,
12};
13
14use crate::{Error, RecvStream, SendStream};
15use web_streams::{Reader, Writer};
16
17struct SharedReader<T: JsCast> {
18 inner: Rc<Mutex<Reader<T>>>,
19}
20
21impl<T: JsCast> SharedReader<T> {
22 fn new(stream: &web_sys::ReadableStream) -> Result<Self, web_streams::Error> {
23 Ok(Self {
24 inner: Rc::new(Mutex::new(Reader::new(stream)?)),
25 })
26 }
27
28 async fn read(&self) -> Result<Option<T>, web_streams::Error> {
29 self.inner.lock().await.read().await
30 }
31}
32
33impl<T: JsCast> Clone for SharedReader<T> {
34 fn clone(&self) -> Self {
35 Self {
36 inner: self.inner.clone(),
37 }
38 }
39}
40
41#[derive(Clone)]
53pub struct Session {
54 inner: WebTransport,
55 url: Url,
56 incoming_uni: SharedReader<WebTransportReceiveStream>,
57 incoming_bi: SharedReader<WebTransportBidirectionalStream>,
58 datagram_reader: SharedReader<Uint8Array>,
59 datagram_writer: Rc<Mutex<Writer>>,
60}
61
62impl Session {
63 pub fn new(inner: WebTransport, url: Url) -> Result<Self, Error> {
64 let incoming_uni = SharedReader::new(&inner.incoming_unidirectional_streams())?;
65 let incoming_bi = SharedReader::new(&inner.incoming_bidirectional_streams())?;
66 let datagrams = inner.datagrams();
67 let datagram_reader = SharedReader::new(&datagrams.readable())?;
68 let datagram_writer = Writer::new(&datagrams.writable())?;
69
70 Ok(Self {
71 inner,
72 url,
73 incoming_uni,
74 incoming_bi,
75 datagram_reader,
76 datagram_writer: Rc::new(Mutex::new(datagram_writer)),
77 })
78 }
79
80 pub async fn accept_uni(&self) -> Result<RecvStream, Error> {
85 match self.incoming_uni.read().await? {
86 Some(stream) => Ok(RecvStream::new(stream)?),
87 None => Err(self.closed().await),
88 }
89 }
90
91 pub async fn accept_bi(&self) -> Result<(SendStream, RecvStream), Error> {
96 let stream: WebTransportBidirectionalStream = match self.incoming_bi.read().await? {
97 Some(stream) => stream,
98 None => return Err(self.closed().await),
99 };
100
101 let send = SendStream::new(stream.writable())?;
102 let recv = RecvStream::new(stream.readable())?;
103
104 Ok((send, recv))
105 }
106
107 pub async fn open_bi(&self) -> Result<(SendStream, RecvStream), Error> {
109 let stream: WebTransportBidirectionalStream =
110 JsFuture::from(self.inner.create_bidirectional_stream()).await?;
111
112 let send = SendStream::new(stream.writable())?;
113 let recv = RecvStream::new(stream.readable())?;
114
115 Ok((send, recv))
116 }
117
118 pub async fn open_uni(&self) -> Result<SendStream, Error> {
120 let stream: WebTransportSendStream =
121 JsFuture::from(self.inner.create_unidirectional_stream()).await?;
122
123 let send = SendStream::new(stream)?;
124 Ok(send)
125 }
126
127 pub async fn send_datagram(&self, payload: Bytes) -> Result<(), Error> {
129 let mut writer = self.datagram_writer.lock().await;
130 writer.write(&Uint8Array::from(payload.as_ref())).await?;
131 Ok(())
132 }
133
134 pub async fn recv_datagram(&self) -> Result<Bytes, Error> {
136 let data: Uint8Array = match self.datagram_reader.read().await? {
137 Some(data) => data,
138 None => return Err(self.closed().await),
139 };
140 Ok(data.to_vec().into())
141 }
142
143 pub fn close(&self, code: u32, reason: &str) {
145 let info = WebTransportCloseInfo::new();
146 info.set_close_code(code);
147 info.set_reason(reason);
148 self.inner.close_with_close_info(&info);
149 }
150
151 pub async fn closed(&self) -> Error {
153 match JsFuture::from(self.inner.closed()).await {
154 Ok(info) => {
155 let info: WebTransportCloseInfo = info;
156 Error::SessionClosed {
157 code: info.get_close_code().unwrap_or_default(),
158 reason: info.get_reason().unwrap_or_default(),
159 }
160 }
161 Err(error) => error.into(),
162 }
163 }
164
165 pub fn url(&self) -> &Url {
167 &self.url
168 }
169}
170
171impl PartialEq for Session {
172 fn eq(&self, other: &Self) -> bool {
173 self.inner == other.inner
174 }
175}
176
177impl Eq for Session {}
178
179#[cfg(target_family = "wasm")]
180impl webtrans_trait::Session for Session {
181 type SendStream = SendStream;
182 type RecvStream = RecvStream;
183 type Error = Error;
184
185 async fn accept_uni(&self) -> Result<Self::RecvStream, Self::Error> {
186 Self::accept_uni(self).await
187 }
188
189 async fn accept_bi(&self) -> Result<(Self::SendStream, Self::RecvStream), Self::Error> {
190 Self::accept_bi(self).await
191 }
192
193 async fn open_bi(&self) -> Result<(Self::SendStream, Self::RecvStream), Self::Error> {
194 Self::open_bi(self).await
195 }
196
197 async fn open_uni(&self) -> Result<Self::SendStream, Self::Error> {
198 Self::open_uni(self).await
199 }
200
201 async fn send_datagram(&self, payload: Bytes) -> Result<(), Self::Error> {
202 Self::send_datagram(self, payload).await
203 }
204
205 async fn recv_datagram(&self) -> Result<Bytes, Self::Error> {
206 Self::recv_datagram(self).await
207 }
208
209 fn max_datagram_size(&self) -> usize {
210 self.inner.datagrams().max_datagram_size() as usize
211 }
212
213 fn close(&self, code: u32, reason: &str) {
214 Self::close(self, code, reason);
215 }
216
217 async fn closed(&self) -> Self::Error {
218 Self::closed(self).await
219 }
220}
221
222#[cfg(all(test, target_family = "wasm"))]
223mod tests {
224 use std::{cell::RefCell, rc::Rc};
225
226 use futures::{join, pin_mut, poll};
227 use js_sys::{Object, Reflect, Uint8Array};
228 use wasm_bindgen::{JsValue, closure::Closure};
229 use wasm_bindgen_test::*;
230 use web_sys::{ReadableStream, ReadableStreamDefaultController};
231
232 use super::SharedReader;
233
234 wasm_bindgen_test_configure!(run_in_browser);
235
236 fn controlled_stream() -> (ReadableStream, ReadableStreamDefaultController) {
237 let controller = Rc::new(RefCell::new(None));
238 let captured = controller.clone();
239 let start = Closure::<dyn FnMut(ReadableStreamDefaultController)>::new(move |value| {
240 *captured.borrow_mut() = Some(value);
241 });
242 let source = Object::new();
243 Reflect::set(&source, &JsValue::from_str("start"), start.as_ref()).unwrap();
244 let stream = ReadableStream::new_with_underlying_source(&source).unwrap();
245 let controller = controller.borrow_mut().take().unwrap();
246 (stream, controller)
247 }
248
249 #[wasm_bindgen_test(async)]
250 async fn cancelled_read_is_resumed_by_the_next_caller() {
251 let (stream, controller) = controlled_stream();
252 let reader = SharedReader::<Uint8Array>::new(&stream).unwrap();
253
254 {
255 let pending = reader.read();
256 pin_mut!(pending);
257 assert!(poll!(pending.as_mut()).is_pending());
258 }
259
260 controller
261 .enqueue_with_chunk(&Uint8Array::from(&b"resumed"[..]).into())
262 .unwrap();
263 let value = reader.read().await.unwrap().unwrap();
264 assert_eq!(value.to_vec(), b"resumed");
265 }
266
267 #[wasm_bindgen_test(async)]
268 async fn cloned_readers_serialize_without_losing_values() {
269 let (stream, controller) = controlled_stream();
270 let first = SharedReader::<Uint8Array>::new(&stream).unwrap();
271 let second = first.clone();
272 controller
273 .enqueue_with_chunk(&Uint8Array::from(&b"one"[..]).into())
274 .unwrap();
275 controller
276 .enqueue_with_chunk(&Uint8Array::from(&b"two"[..]).into())
277 .unwrap();
278
279 let (one, two) = join!(first.read(), second.read());
280 assert_eq!(one.unwrap().unwrap().to_vec(), b"one");
281 assert_eq!(two.unwrap().unwrap().to_vec(), b"two");
282 }
283}