diff --git a/client/src/pipe.rs b/client/src/pipe.rs index 136fe01..8e839dd 100644 --- a/client/src/pipe.rs +++ b/client/src/pipe.rs @@ -78,41 +78,9 @@ pub struct PipeRequest { pub(crate) pipe_id: u32, pub(crate) description: String, pub(crate) sender: Sender, - pub(crate) receiver: Receiver, pub(crate) dispatcher: Arc, } -#[cfg(feature = "pipes")] -struct ExpectedPipeGuard { - receiver: Receiver, - pipe_id: u32, - armed: bool, -} - -#[cfg(feature = "pipes")] -impl ExpectedPipeGuard { - fn new(receiver: Receiver, pipe_id: u32) -> Self { - Self { - receiver, - pipe_id, - armed: true, - } - } - - fn disarm(&mut self) { - self.armed = false; - } -} - -#[cfg(feature = "pipes")] -impl Drop for ExpectedPipeGuard { - fn drop(&mut self) { - if self.armed { - self.receiver.cancel_expected_pipe(self.pipe_id); - } - } -} - #[cfg(feature = "pipes")] impl PipeRequest { pub fn id(&self) -> u32 { @@ -124,10 +92,6 @@ impl PipeRequest { } pub async fn accept(self) -> Result { - self.receiver - .expect_pipe(self.pipe_id) - .map_err(PipeError::from)?; - let mut expected_pipe = ExpectedPipeGuard::new(self.receiver.clone(), self.pipe_id); let (pipe_tx, pipe_rx) = tokio::sync::oneshot::channel(); { let mut pending = self.dispatcher.pending_pipes.lock().await; @@ -151,10 +115,7 @@ impl PipeRequest { let timeout = self.dispatcher.policy.read_timeout; match tokio::time::timeout(timeout, pipe_rx).await { - Ok(Ok(reader)) => { - expected_pipe.disarm(); - Ok(reader) - } + Ok(Ok(reader)) => Ok(reader), Ok(Err(_)) => { self.dispatcher .pending_pipes @@ -452,7 +413,6 @@ pub(crate) async fn run_dispatcher( pipe_id, description, sender: sender.clone(), - receiver: receiver.clone(), dispatcher: dispatcher.clone(), }; let _ = pipe_req_tx.send(req).await; diff --git a/common/src/lib.rs b/common/src/lib.rs index 71e87f2..1a7fe66 100644 --- a/common/src/lib.rs +++ b/common/src/lib.rs @@ -164,9 +164,6 @@ pub enum CommunicationError { #[error("Stream Error")] StreamError, - #[error("Stream failed after delivery may have started")] - DeliveryUnknown, - #[error("Stream Error: {0}")] #[cfg(not(target_arch = "wasm32"))] StreamWriteError(#[from] wtransport::error::StreamWriteError), @@ -185,38 +182,6 @@ pub enum CommunicationError { Other(String), } -/// How the protocol layer should handle the first frame on a receive stream. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum FirstFrameDisposition { - Message, - Pipe(u32), -} - -/// Classify a first frame without tying the decision to a WebTransport backend. -/// -/// `PipeRequest` is used both as a control message and as the header of the raw -/// stream opened after that request is accepted. Only the protocol layer knows -/// which raw stream IDs are currently expected. -pub fn classify_first_frame( - is_pipe_request: bool, - pipe_id: Option, - pipe_is_expected: bool, -) -> Result { - if !is_pipe_request { - return Ok(FirstFrameDisposition::Message); - } - - let pipe_id = pipe_id.filter(|id| *id != 0).ok_or_else(|| { - CommunicationError::Other("PipeRequest frame must contain a non-zero id".into()) - })?; - - if pipe_is_expected { - Ok(FirstFrameDisposition::Pipe(pipe_id)) - } else { - Ok(FirstFrameDisposition::Message) - } -} - // ---- manual PartialEq (quinn / wtransport types don't impl PartialEq) ---- impl PartialEq for CommunicationError { @@ -247,7 +212,6 @@ impl PartialEq for CommunicationError { (Self::ReadExactError(_), Self::ReadExactError(_)) => true, (Self::StreamClosed, Self::StreamClosed) => true, (Self::StreamError, Self::StreamError) => true, - (Self::DeliveryUnknown, Self::DeliveryUnknown) => true, #[cfg(not(target_arch = "wasm32"))] (Self::StreamWriteError(_), Self::StreamWriteError(_)) => true, #[cfg(not(target_arch = "wasm32"))] diff --git a/host/src/connection.rs b/host/src/connection.rs index 7871d55..71c60db 100644 --- a/host/src/connection.rs +++ b/host/src/connection.rs @@ -73,7 +73,7 @@ pub struct MTPConnection< #[cfg(feature = "pipes")] pub(crate) app_rx: Mutex>>, #[cfg(feature = "pipes")] - pub(crate) pipe_req_rx: Mutex>>, + pub(crate) pipe_req_rx: Mutex>>, #[cfg(feature = "pipes")] pub(crate) pipe_dispatcher: Arc>, #[cfg(not(feature = "pipes"))] @@ -381,7 +381,7 @@ where }) } - pub async fn receive_pipe(&self) -> Result, CommunicationError> { + pub async fn receive_pipe(&self) -> Result, CommunicationError> { self.pipe_req_rx .lock() .await diff --git a/host/src/pipe.rs b/host/src/pipe.rs index 192e3d7..eae1383 100644 --- a/host/src/pipe.rs +++ b/host/src/pipe.rs @@ -27,10 +27,6 @@ pub trait PipeReceiver

: Clone + Send + Sync + 'static where P: tokio::io::AsyncRead + Send + Unpin + 'static, { - fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError>; - - fn cancel_expected_pipe(&self, pipe_id: u32); - fn receive_pipe_event( &self, ) -> impl std::future::Future, CommunicationError>> + Send; @@ -56,14 +52,6 @@ impl PipeSender for mtp_transport::Sender { } impl PipeReceiver for mtp_transport::Receiver { - fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError> { - self.expect_pipe(pipe_id) - } - - fn cancel_expected_pipe(&self, pipe_id: u32) { - self.cancel_expected_pipe(pipe_id); - } - async fn receive_pipe_event( &self, ) -> Result, CommunicationError> { @@ -99,14 +87,6 @@ where C: mtp_transport::TransportConnection, C::RecvStream: tokio::io::AsyncRead + Send + Unpin + 'static, { - fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError> { - self.expect_pipe(pipe_id) - } - - fn cancel_expected_pipe(&self, pipe_id: u32) { - self.cancel_expected_pipe(pipe_id); - } - async fn receive_pipe_event( &self, ) -> Result, CommunicationError> { @@ -172,60 +152,16 @@ where } } -pub struct PipeRequest { +pub struct PipeRequest { pub(crate) pipe_id: u32, pub(crate) description: String, pub(crate) sender: S, - pub(crate) receiver: R, pub(crate) dispatcher: Arc>, } -struct ExpectedPipeGuard -where - R: PipeReceiver

, - P: tokio::io::AsyncRead + Send + Unpin + 'static, -{ - receiver: R, - pipe_id: u32, - armed: bool, - _stream: std::marker::PhantomData

, -} - -impl ExpectedPipeGuard -where - R: PipeReceiver

, - P: tokio::io::AsyncRead + Send + Unpin + 'static, -{ - fn new(receiver: R, pipe_id: u32) -> Self { - Self { - receiver, - pipe_id, - armed: true, - _stream: std::marker::PhantomData, - } - } - - fn disarm(&mut self) { - self.armed = false; - } -} - -impl Drop for ExpectedPipeGuard -where - R: PipeReceiver

, - P: tokio::io::AsyncRead + Send + Unpin + 'static, -{ - fn drop(&mut self) { - if self.armed { - self.receiver.cancel_expected_pipe(self.pipe_id); - } - } -} - -impl PipeRequest +impl PipeRequest where S: PipeSender, - R: PipeReceiver

, P: tokio::io::AsyncRead + Send + Unpin + 'static, { pub fn id(&self) -> u32 { @@ -237,10 +173,6 @@ where } pub async fn accept(self) -> Result, PipeError> { - self.receiver - .expect_pipe(self.pipe_id) - .map_err(PipeError::from)?; - let mut expected_pipe = ExpectedPipeGuard::::new(self.receiver.clone(), self.pipe_id); let (pipe_tx, pipe_rx) = tokio::sync::oneshot::channel(); self.dispatcher .pending_pipes @@ -264,10 +196,7 @@ where } match tokio::time::timeout(self.dispatcher.policy.read_timeout, pipe_rx).await { - Ok(Ok(reader)) => { - expected_pipe.disarm(); - Ok(reader) - } + Ok(Ok(reader)) => Ok(reader), Ok(Err(_)) => { self.dispatcher .pending_pipes @@ -432,7 +361,7 @@ pub(crate) async fn run_dispatcher( receiver: R, sender: S, app_tx: mpsc::Sender>, - pipe_req_tx: mpsc::Sender>, + pipe_req_tx: mpsc::Sender>, dispatcher: Arc>, ) where S: PipeSender, @@ -459,7 +388,6 @@ pub(crate) async fn run_dispatcher( .unwrap_or("") .to_owned(), sender: sender.clone(), - receiver: receiver.clone(), dispatcher: dispatcher.clone(), }; let _ = pipe_req_tx.send(request).await; @@ -518,7 +446,6 @@ pub(crate) async fn run_dispatcher( pipe_id, description: reader.description().to_owned(), sender: sender.clone(), - receiver: receiver.clone(), dispatcher: dispatcher.clone(), }; let _ = pipe_req_tx.send(request).await; diff --git a/mtp-webserver/src/transport.rs b/mtp-webserver/src/transport.rs index 9dcce5f..9b7de76 100644 --- a/mtp-webserver/src/transport.rs +++ b/mtp-webserver/src/transport.rs @@ -57,14 +57,14 @@ impl TransportSendStream for H3TransportSender { self.stream .write_all(buf) .await - .map_err(|_| CommunicationError::DeliveryUnknown)?; + .map_err(|_| CommunicationError::StreamError)?; // Control/authentication frames use a persistent stream. h3 keeps // those writes buffered until flushed; without this the peer can wait // for the challenge while the server waits for its proof. self.stream .flush() .await - .map_err(|_| CommunicationError::DeliveryUnknown) + .map_err(|_| CommunicationError::StreamError) } async fn finish(&mut self) -> Result<(), CommunicationError> { @@ -73,11 +73,6 @@ impl TransportSendStream for H3TransportSender { .await .map_err(|_| CommunicationError::StreamError) } - - fn reset(&mut self, code: u32) -> Result<(), CommunicationError> { - h3::quic::SendStream::reset(&mut self.stream, code as u64); - Ok(()) - } } #[async_trait::async_trait] @@ -145,11 +140,6 @@ impl TransportRecvStream for H3TransportReceiver { } } } - - fn stop(mut self, code: u32) -> Result<(), CommunicationError> { - h3::quic::RecvStream::stop_sending(&mut self.stream, code as u64); - Ok(()) - } } impl tokio::io::AsyncWrite for H3TransportSender { diff --git a/transport/src/connection.rs b/transport/src/connection.rs index 7db09e8..dece604 100644 --- a/transport/src/connection.rs +++ b/transport/src/connection.rs @@ -4,10 +4,6 @@ use crate::framing::RetryClassifier; use crate::pipe::PipeReader; use mtp_codec::{CommunicationValue, DecodeError, DecodeLimits, EncodeLimits, TypeMap}; use mtp_common::CommunicationError; -#[cfg(feature = "pipes")] -use mtp_common::{FirstFrameDisposition, classify_first_frame}; -#[cfg(feature = "pipes")] -use std::collections::HashSet; use std::ops::Deref; use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; @@ -292,15 +288,15 @@ impl Sender { Ok(Ok(())) => Ok(()), Ok(Err(wtransport::error::StreamWriteError::Stopped(code))) => { warn!("[Sender] write failed: peer sent STOP_SENDING (error code {code})"); - Err(CommunicationError::DeliveryUnknown) + Err(CommunicationError::StreamClosed) } Ok(Err(other)) => { warn!("[Sender] write failed: {other}"); - Err(CommunicationError::DeliveryUnknown) + Err(CommunicationError::StreamError) } Err(_) => { warn!("[Sender] write timed out (len={})", bytes.len()); - Err(CommunicationError::DeliveryUnknown) + Err(CommunicationError::StreamError) } } } @@ -402,15 +398,15 @@ impl Sender { Ok(Ok(())) => Ok(()), Ok(Err(wtransport::error::StreamWriteError::Stopped(code))) => { warn!("[Sender] finish failed: peer sent STOP_SENDING (error code {code})"); - Err(CommunicationError::DeliveryUnknown) + Err(CommunicationError::StreamClosed) } Ok(Err(other)) => { warn!("[Sender] finish failed: {other}"); - Err(CommunicationError::DeliveryUnknown) + Err(CommunicationError::StreamError) } Err(_) => { warn!("[Sender] finish timed out"); - Err(CommunicationError::DeliveryUnknown) + Err(CommunicationError::StreamError) } } } @@ -749,8 +745,6 @@ struct ReceiverInner { max_message_size: Arc, type_map: Arc>, decode_rejections: Arc, - #[cfg(feature = "pipes")] - expected_pipes: Arc>>, } impl Clone for Receiver { @@ -837,10 +831,6 @@ impl Receiver { let accept_type_map = type_map.clone(); let decode_rejections = Arc::new(DecodeRejectionCounters::default()); let accept_decode_rejections = decode_rejections.clone(); - #[cfg(feature = "pipes")] - let expected_pipes = Arc::new(std::sync::Mutex::new(HashSet::new())); - #[cfg(feature = "pipes")] - let accept_expected_pipes = expected_pipes.clone(); let stream_limit = Arc::new(Semaphore::new(policy.max_concurrent_stream_tasks.max(1))); let accept_stream_limit = stream_limit.clone(); debug!( @@ -910,8 +900,6 @@ impl Receiver { let stream_max_message_size = accept_max_message_size.clone(); let stream_type_map = accept_type_map.clone(); let stream_decode_rejections = accept_decode_rejections.clone(); - #[cfg(feature = "pipes")] - let stream_expected_pipes = accept_expected_pipes.clone(); tokio::spawn(async move { let _permit = permit; @@ -947,54 +935,38 @@ impl Receiver { #[cfg(feature = "pipes")] { - if frame_count == 1 { - let is_pipe_request = msg.is_type( - mtp_codec::CommunicationType::PipeRequest, - ); - let pipe_id = msg.id().filter(|id| *id != 0); - let pipe_is_expected = is_pipe_request && pipe_id.is_some_and(|pipe_id| { - stream_expected_pipes - .lock() - .is_ok_and(|mut expected| expected.remove(&pipe_id)) - }); - let disposition = match classify_first_frame( - is_pipe_request, - msg.id(), - pipe_is_expected, - ) { - Ok(disposition) => disposition, - Err(error) => { - let _ = msg_tx_stream - .send(Err(error.clone())) - .await; - stream_handle.close(Some(error)); - break; - } + if msg.is_type(mtp_codec::CommunicationType::PipeRequest) + && frame_count == 1 + { + let Some(pipe_id) = msg.id().filter(|id| *id != 0) else { + let error = CommunicationError::Other( + "PipeRequest frame must contain a non-zero id".into(), + ); + let _ = msg_tx_stream.send(Err(error.clone())).await; + stream_handle.close(Some(error)); + break; + }; + let description = msg + .get_str(mtp_codec::DataType::Description) + .unwrap_or("") + .to_string(); + + let pipe_reader = crate::pipe::PipeReader { + stream: s, + description, + pipe_id, }; - if let FirstFrameDisposition::Pipe(pipe_id) = disposition { - let description = msg - .get_str(mtp_codec::DataType::Description) - .unwrap_or("") - .to_string(); - - let pipe_reader = crate::pipe::PipeReader { - stream: s, - description, - pipe_id, - }; - - if pipe_tx_stream - .send(pipe_reader) - .await - .is_err() - { - stream_handle.close(Some( - CommunicationError::StreamClosed, - )); - } - break; + if pipe_tx_stream + .send(pipe_reader) + .await + .is_err() + { + stream_handle.close(Some( + CommunicationError::StreamClosed, + )); } + break; } } @@ -1139,8 +1111,6 @@ impl Receiver { max_message_size, type_map, decode_rejections, - #[cfg(feature = "pipes")] - expected_pipes, }), } } @@ -1157,26 +1127,6 @@ impl Receiver { *self.inner.type_map.write().await = type_map.clone(); } - #[cfg(feature = "pipes")] - pub fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError> { - if pipe_id == 0 { - return Err(CommunicationError::Other("pipe id must be non-zero".into())); - } - self.inner - .expected_pipes - .lock() - .map_err(|_| CommunicationError::Other("expected pipe state is unavailable".into()))? - .insert(pipe_id); - Ok(()) - } - - #[cfg(feature = "pipes")] - pub fn cancel_expected_pipe(&self, pipe_id: u32) { - if let Ok(mut expected) = self.inner.expected_pipes.lock() { - expected.remove(&pipe_id); - } - } - /// Return local counts for frames rejected by the structured decoder. /// /// These counters are intentionally local-only; peers continue to receive diff --git a/transport/src/framing.rs b/transport/src/framing.rs index ac4e005..9fa1d80 100644 --- a/transport/src/framing.rs +++ b/transport/src/framing.rs @@ -82,10 +82,6 @@ mod tests { async fn finish(&mut self) -> Result<(), CommunicationError> { Ok(()) } - - fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> { - Ok(()) - } } #[tokio::test] diff --git a/transport/src/generic_connection.rs b/transport/src/generic_connection.rs index 4361855..a7a32de 100644 --- a/transport/src/generic_connection.rs +++ b/transport/src/generic_connection.rs @@ -10,12 +10,8 @@ use crate::{ framing::{RetryClassifier, write_frame}, }; use mtp_codec::{CommunicationValue, DataType, DecodeLimits, TypeMap}; -use mtp_common::{CommunicationError, FirstFrameDisposition, classify_first_frame}; -#[cfg(feature = "pipes")] -use std::collections::HashSet; +use mtp_common::CommunicationError; use std::sync::Arc; -#[cfg(feature = "pipes")] -use std::sync::Mutex as StdMutex; use std::sync::atomic::{AtomicU64, Ordering}; use tokio::sync::{Mutex, Notify, RwLock, Semaphore, mpsc}; use tokio::time::{Instant, timeout, timeout_at}; @@ -89,10 +85,9 @@ impl GenericSender { ) .await .map_err(|_| CommunicationError::StreamError)??; - match timeout(self.policy.write_timeout, stream.finish()).await { - Ok(Ok(())) => Ok(()), - Ok(Err(_)) | Err(_) => Err(CommunicationError::DeliveryUnknown), - } + timeout(self.policy.write_timeout, stream.finish()) + .await + .map_err(|_| CommunicationError::StreamError)? } crate::SendMode::PersistentStream => { let mut stream = self.persistent.lock().await; @@ -209,8 +204,6 @@ pub struct GenericReceiver { type_map: Arc>, queue_notify: Arc, decode_rejections: Arc, - #[cfg(feature = "pipes")] - expected_pipes: Arc>>, _accept_task: Arc>, } @@ -226,8 +219,6 @@ impl Clone for GenericReceiver { type_map: self.type_map.clone(), queue_notify: self.queue_notify.clone(), decode_rejections: self.decode_rejections.clone(), - #[cfg(feature = "pipes")] - expected_pipes: self.expected_pipes.clone(), _accept_task: self._accept_task.clone(), } } @@ -263,10 +254,6 @@ impl GenericReceiver { let task_queue_notify = queue_notify.clone(); let decode_rejections = Arc::new(DecodeRejectionCounters::default()); let task_decode_rejections = decode_rejections.clone(); - #[cfg(feature = "pipes")] - let expected_pipes = Arc::new(StdMutex::new(HashSet::new())); - #[cfg(feature = "pipes")] - let task_expected_pipes = expected_pipes.clone(); let task_accept_task_tx = tx.clone(); #[cfg(feature = "pipes")] let task_accept_task_pipe_tx = pipe_tx.clone(); @@ -325,8 +312,6 @@ impl GenericReceiver { let connection = task_connection.clone(); let type_map = task_type_map.clone(); let decode_rejections = task_decode_rejections.clone(); - #[cfg(feature = "pipes")] - let expected_pipes = task_expected_pipes.clone(); tokio::spawn(async move { let _permit = permit; let mut stream = stream; @@ -458,51 +443,37 @@ impl GenericReceiver { #[cfg(feature = "pipes")] { - if frames == 1 { - let is_pipe_request = - message.is_type(mtp_codec::CommunicationType::PipeRequest); - let pipe_id = message.id().filter(|id| *id != 0); - let pipe_is_expected = is_pipe_request - && pipe_id.is_some_and(|pipe_id| { - expected_pipes - .lock() - .is_ok_and(|mut expected| expected.remove(&pipe_id)) - }); - let disposition = match classify_first_frame( - is_pipe_request, - message.id(), - pipe_is_expected, - ) { - Ok(disposition) => disposition, - Err(error) => { - let _ = tx.send(Err(error.clone())).await; - connection.close( - policy.application_close_code, - b"pipe request missing id", - ); - break; - } + if message.is_type(mtp_codec::CommunicationType::PipeRequest) + && frames == 1 + { + let Some(pipe_id) = message.id().filter(|id| *id != 0) else { + let error = CommunicationError::Other( + "PipeRequest frame must contain a non-zero id".into(), + ); + let _ = tx.send(Err(error.clone())).await; + connection.close( + policy.application_close_code, + b"pipe request missing id", + ); + break; + }; + let description = message + .get_str(mtp_codec::DataType::Description) + .unwrap_or("") + .to_string(); + + let pipe_reader = PipeReader { + stream, + description, + pipe_id, }; - if let FirstFrameDisposition::Pipe(pipe_id) = disposition { - let description = message - .get_str(mtp_codec::DataType::Description) - .unwrap_or("") - .to_string(); + tracing::debug!(pipe_id, description = %pipe_reader.description, "classified incoming pipe stream"); - let pipe_reader = PipeReader { - stream, - description, - pipe_id, - }; - - tracing::debug!(pipe_id, description = %pipe_reader.description, "classified incoming pipe stream"); - - if pipe_tx.send(pipe_reader).await.is_err() { - break; - } - return; + if pipe_tx.send(pipe_reader).await.is_err() { + break; } + return; } } @@ -546,8 +517,6 @@ impl GenericReceiver { type_map, queue_notify, decode_rejections, - #[cfg(feature = "pipes")] - expected_pipes, _accept_task: Arc::new(accept_task), } } @@ -555,25 +524,6 @@ impl GenericReceiver { *self.ping_sender.write().await = Some(sender); } - #[cfg(feature = "pipes")] - pub fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError> { - if pipe_id == 0 { - return Err(CommunicationError::Other("pipe id must be non-zero".into())); - } - self.expected_pipes - .lock() - .map_err(|_| CommunicationError::Other("expected pipe state is unavailable".into()))? - .insert(pipe_id); - Ok(()) - } - - #[cfg(feature = "pipes")] - pub fn cancel_expected_pipe(&self, pipe_id: u32) { - if let Ok(mut expected) = self.expected_pipes.lock() { - expected.remove(&pipe_id); - } - } - /// Switch from the handshake frame limit to the application frame limit. pub fn set_max_message_size(&self, max_message_size: u64) { self.max_message_size diff --git a/transport/src/transport_traits.rs b/transport/src/transport_traits.rs index 70395af..63c0268 100644 --- a/transport/src/transport_traits.rs +++ b/transport/src/transport_traits.rs @@ -17,7 +17,6 @@ use mtp_common::CommunicationError; pub trait TransportSendStream: tokio::io::AsyncWrite + Send + Sync { async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError>; async fn finish(&mut self) -> Result<(), CommunicationError>; - fn reset(&mut self, code: u32) -> Result<(), CommunicationError>; } /// A readable unidirectional stream suitable for MTP frames. @@ -29,9 +28,6 @@ pub trait TransportSendStream: tokio::io::AsyncWrite + Send + Sync { pub trait TransportRecvStream: tokio::io::AsyncRead + Send + Sync { async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError>; async fn read_chunk(&mut self, max: usize) -> Result>, CommunicationError>; - fn stop(self, code: u32) -> Result<(), CommunicationError> - where - Self: Sized; } /// A QUIC/WebTransport connection that provides MTP's unidirectional streams. @@ -51,7 +47,7 @@ impl TransportSendStream for wtransport::SendStream { async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError> { wtransport::SendStream::write_all(self, buf) .await - .map_err(|_| CommunicationError::DeliveryUnknown) + .map_err(|_| CommunicationError::StreamError) } async fn finish(&mut self) -> Result<(), CommunicationError> { @@ -59,23 +55,14 @@ impl TransportSendStream for wtransport::SendStream { .await .map_err(|_| CommunicationError::StreamError) } - - fn reset(&mut self, code: u32) -> Result<(), CommunicationError> { - wtransport::SendStream::reset(self, wtransport::VarInt::from_u32(code)) - .map_err(|_| CommunicationError::StreamClosed) - } } #[async_trait] impl TransportRecvStream for wtransport::RecvStream { async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError> { - match wtransport::RecvStream::read_exact(self, buf).await { - Ok(()) => Ok(()), - Err(wtransport::error::StreamReadExactError::FinishedEarly(0)) => { - Err(CommunicationError::StreamClosed) - } - Err(_) => Err(CommunicationError::StreamError), - } + wtransport::RecvStream::read_exact(self, buf) + .await + .map_err(|_| CommunicationError::StreamError) } async fn read_chunk(&mut self, max: usize) -> Result>, CommunicationError> { @@ -89,11 +76,6 @@ impl TransportRecvStream for wtransport::RecvStream { Err(_) => Err(CommunicationError::StreamError), } } - - fn stop(self, code: u32) -> Result<(), CommunicationError> { - wtransport::RecvStream::stop(self, wtransport::VarInt::from_u32(code)); - Ok(()) - } } #[async_trait] diff --git a/transport/tests/generic_pipe.rs b/transport/tests/generic_pipe.rs index 7dfb538..81a875f 100644 --- a/transport/tests/generic_pipe.rs +++ b/transport/tests/generic_pipe.rs @@ -4,7 +4,7 @@ use async_trait::async_trait; use mtp_codec::CommunicationValue; use mtp_common::CommunicationError; use mtp_transport::{ - GenericReceiver, GenericSender, Policy, SendMode, TransportConnection, TransportEvent, + GenericReceiver, GenericSender, Policy, TransportConnection, TransportEvent, TransportRecvStream, TransportSendStream, }; use std::sync::Arc; @@ -53,10 +53,6 @@ impl TransportSendStream for MockSendStream { .await .map_err(|_| CommunicationError::StreamError) } - - fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> { - Ok(()) - } } struct MockRecvStream { @@ -93,10 +89,6 @@ impl TransportRecvStream for MockRecvStream { Err(_) => Err(CommunicationError::StreamError), } } - - fn stop(self, _code: u32) -> Result<(), CommunicationError> { - Ok(()) - } } #[derive(Clone)] @@ -165,7 +157,6 @@ async fn test_open_pipe_and_receive_reader() -> Result<(), Box Result<(), Box let sender = GenericSender::new(conn_a, policy.clone()); let receiver = GenericReceiver::new(conn_b, policy); - receiver.expect_pipe(1)?; let mut pipe_writer = sender.open_pipe(1, "data-pipe").await?; let data = b"hello through the pipe"; @@ -205,7 +195,6 @@ async fn test_pipe_large_payload() -> Result<(), Box> { let sender = GenericSender::new(conn_a, policy.clone()); let receiver = GenericReceiver::new(conn_b, policy); - receiver.expect_pipe(7)?; let mut pipe_writer = sender.open_pipe(7, "big-pipe").await?; let data: Vec = (0..256 * 1024).map(|i| (i % 256) as u8).collect(); @@ -234,7 +223,6 @@ async fn test_receive_event_dispatches_pipe() -> Result<(), Box Result<(), Box Result<(), Box> { let (conn_a, conn_b) = mock_connected_pair().await; - let policy = Arc::new(Policy::default().with_send_mode(SendMode::SingleStreamPerMessage)); + let policy = Arc::new(Policy::default()); let sender = GenericSender::new(conn_a, policy.clone()); let receiver = GenericReceiver::new(conn_b, policy); - let request = CommunicationValue::new(mtp_codec::CommunicationType::PipeRequest) - .with_id(1) - .add_typed_default( - mtp_codec::DataType::Description, - mtp_codec::DataValue::Str("mixed-pipe".into()), - ); - sender.send(&request).await?; + let msg = CommunicationValue::new(mtp_codec::CommunicationType::BadRequest); + sender.send(&msg).await?; + + let _pipe_writer = sender.open_pipe(1, "mixed-pipe").await?; let received = receiver.receive().await?; - assert!(received.is_type(mtp_codec::CommunicationType::PipeRequest)); - - receiver.expect_pipe(1)?; - let _pipe_writer = sender.open_pipe(1, "mixed-pipe").await?; + assert_eq!( + received.get_type(), + mtp_codec::CommunicationType::BadRequest + .try_to_id(&mtp_codec::TypeMap::latest()) + .unwrap() + ); let pipe_reader = receiver.receive_pipe().await?; assert_eq!(pipe_reader.pipe_id(), 1); @@ -304,8 +291,6 @@ async fn test_multiple_pipes() -> Result<(), Box> { let sender = GenericSender::new(conn_a, policy.clone()); let receiver = GenericReceiver::new(conn_b, policy); - receiver.expect_pipe(10)?; - receiver.expect_pipe(20)?; let mut pw1 = sender.open_pipe(10, "first").await?; let mut pw2 = sender.open_pipe(20, "second").await?; diff --git a/wasm/src/client/receive.rs b/wasm/src/client/receive.rs index 51a48bf..35ceb02 100644 --- a/wasm/src/client/receive.rs +++ b/wasm/src/client/receive.rs @@ -179,7 +179,6 @@ impl WasmClient { let expired_pipe_creations = self.expired_pipe_creations.clone(); let pending_pipes = self.pending_pipes.clone(); let loop_pending_pipes = pending_pipes.clone(); - let expected_pending_pipes = pending_pipes.clone(); let on_pipe_request = self.on_pipe_request.clone(); let loop_pipe_creations = pending_pipe_creations.clone(); let loop_expired_pipe_creations = expired_pipe_creations.clone(); @@ -294,12 +293,6 @@ impl WasmClient { let _ = entry.sender.send(Ok(pipe_reader)); } }, - move |pipe_id| { - expected_pending_pipes - .borrow() - .get(&pipe_id) - .is_some_and(|entry| entry.generation == loop_generation) - }, ) .await; if connection_generation.get() != generation { diff --git a/wasm/src/pipe.rs b/wasm/src/pipe.rs index 88783ae..efa4f07 100644 --- a/wasm/src/pipe.rs +++ b/wasm/src/pipe.rs @@ -1,6 +1,9 @@ +use wasm_bindgen::JsCast; use wasm_bindgen::prelude::*; +use wasm_bindgen_futures::JsFuture; -use crate::transport::{BrowserRecvStream, BrowserSendStream, log_stream_error_code}; +use crate::error::js_error; +use crate::transport::release_writer_lock; #[wasm_bindgen(typescript_custom_section)] const PIPE_TS: &str = r#" @@ -20,41 +23,54 @@ export interface PipeReader { #[wasm_bindgen] pub struct PipeWriter { - stream: BrowserSendStream, + writer: JsValue, pipe_id: u32, } impl PipeWriter { - pub(crate) fn new(stream: BrowserSendStream, pipe_id: u32) -> Self { - Self { stream, pipe_id } - } -} - -impl Drop for PipeWriter { - fn drop(&mut self) { - self.stream.release(); + pub fn new(writer: JsValue, pipe_id: u32) -> Self { + Self { writer, pipe_id } } } #[wasm_bindgen] impl PipeWriter { pub async fn write(&mut self, data: &[u8]) -> Result<(), JsValue> { - self.stream.write_all(data).await + let chunk = js_sys::Uint8Array::from(data); + let write_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("write")) + .map_err(|_| js_error("missing write"))? + .dyn_into::() + .map_err(|_| js_error("write not a function"))?; + let write_promise = write_fn + .call1(&self.writer, &chunk) + .map_err(|e| js_error(format!("write failed: {:?}", e)))?; + JsFuture::from(write_promise.unchecked_into::()).await?; + Ok(()) } - pub async fn close(mut self) -> Result<(), JsValue> { - let result = self.stream.finish().await; - if let Err(error) = &result { - log_stream_error_code(error, "pipe writer close"); + pub async fn close(self) -> Result<(), JsValue> { + let close_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("close")) + .map_err(|_| js_error("missing close"))? + .dyn_into::() + .map_err(|_| js_error("close not a function"))?; + let close_promise = close_fn + .call0(&self.writer) + .map_err(|e| js_error(format!("close failed: {:?}", e)))?; + if let Err(e) = JsFuture::from(close_promise.unchecked_into::()).await { + crate::transport::log_stream_error_code(&e, "pipe writer close"); } - self.stream.release(); - result + release_writer_lock(&self.writer); + Ok(()) } pub fn abort(&mut self) -> Result<(), JsValue> { - let result = self.stream.reset(0); - self.stream.release(); - result + let abort_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("abort")) + .map_err(|_| js_error("missing abort"))? + .dyn_into::() + .map_err(|_| js_error("abort not a function"))?; + let _ = abort_fn.call0(&self.writer); + release_writer_lock(&self.writer); + Ok(()) } pub fn pipe_id(&self) -> u32 { @@ -64,26 +80,19 @@ impl PipeWriter { #[wasm_bindgen] pub struct PipeReader { - stream: BrowserRecvStream, + reader: JsValue, description: String, pipe_id: u32, pending: Vec, - finished: bool, } impl PipeReader { - pub(crate) fn new( - stream: BrowserRecvStream, - pipe_id: u32, - description: String, - pending: Vec, - ) -> Self { + pub fn new(reader: JsValue, pipe_id: u32, description: String, pending: Vec) -> Self { Self { - stream, + reader, pipe_id, description, pending, - finished: false, } } } @@ -96,18 +105,27 @@ impl PipeReader { return Ok(js_sys::Uint8Array::from(&data[..]).into()); } - if self.finished { + let read_fn = js_sys::Reflect::get(&self.reader, &JsValue::from_str("read")) + .map_err(|_| js_error("missing read"))? + .dyn_into::() + .map_err(|_| js_error("read not a function"))?; + let promise = read_fn + .call0(&self.reader) + .map_err(|_| js_error("read call failed"))? + .unchecked_into::(); + let result = JsFuture::from(promise).await?; + + let done = js_sys::Reflect::get(&result, &JsValue::from_str("done")) + .ok() + .and_then(|v| v.as_bool()) + .unwrap_or(true); + if done { return Ok(JsValue::NULL); } - match self.stream.read_chunk().await? { - Some(value) => Ok(js_sys::Uint8Array::from(&value[..]).into()), - None => { - self.stream.release(); - self.finished = true; - Ok(JsValue::NULL) - } - } + let value = js_sys::Reflect::get(&result, &JsValue::from_str("value")) + .map_err(|_| js_error("missing value"))?; + Ok(js_sys::Uint8Array::new(&value).into()) } pub fn pipe_id(&self) -> u32 { @@ -118,9 +136,3 @@ impl PipeReader { self.description.clone() } } - -impl Drop for PipeReader { - fn drop(&mut self) { - self.stream.release(); - } -} diff --git a/wasm/src/transport.rs b/wasm/src/transport.rs index b7ffebe..dc0091e 100644 --- a/wasm/src/transport.rs +++ b/wasm/src/transport.rs @@ -9,7 +9,6 @@ use wasm_bindgen_futures::JsFuture; use crate::error::js_error; use crate::frame::parse_frame_value_with_limits; use mtp_codec::{DecodeLimits, EncodeLimits, TypeMap}; -use mtp_common::{FirstFrameDisposition, classify_first_frame}; const CLOSE_FRAME_LEN: u32 = u32::MAX; @@ -26,6 +25,12 @@ pub(crate) fn log_stream_error_code(error: &JsValue, context: &str) { let stream_error_code = js_sys::Reflect::get(error, &JsValue::from_str("streamErrorCode")) .ok() .and_then(|v| v.as_f64()); + if matches!(stream_error_code, Some(0.0)) { + // WebTransport reports peer-driven stream shutdown as code 0 in this + // environment. For one-frame handshake streams, that is expected and + // should not be surfaced as a warning. + return; + } let message = error .as_string() .or_else(|| { @@ -71,233 +76,6 @@ fn resolve_stream_readable(recv_stream: &JsValue) -> Result { } } -#[derive(Clone)] -struct BrowserConnection { - inner: JsValue, - incoming_reader: Rc>>, -} - -pub(crate) struct BrowserSendStream { - writer: JsValue, -} - -pub(crate) struct BrowserRecvStream { - reader: JsValue, -} - -impl BrowserConnection { - async fn connect(url: &str, cert_hashes: Option>) -> Result { - let constructor = - js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("WebTransport"))? - .dyn_into::() - .map_err(|_| js_error("WebTransport not available"))?; - let args = js_sys::Array::new(); - args.push(&JsValue::from_str(url)); - - if let Some(hashes) = cert_hashes { - let webtransport_hashes = js_sys::Array::new(); - for hash in hashes { - let (algorithm, value) = hash.split_once(':').unwrap_or(("sha-256", hash.as_str())); - if let Ok(value) = hex::decode(value) { - let entry = js_sys::Object::new(); - js_sys::Reflect::set( - &entry, - &JsValue::from_str("algorithm"), - &JsValue::from_str(algorithm), - )?; - js_sys::Reflect::set( - &entry, - &JsValue::from_str("value"), - &js_sys::Uint8Array::from(&value[..]), - )?; - webtransport_hashes.push(&entry); - } - } - if webtransport_hashes.length() > 0 { - let options = js_sys::Object::new(); - js_sys::Reflect::set( - &options, - &JsValue::from_str("serverCertificateHashes"), - &webtransport_hashes, - )?; - args.push(&options); - } - } - - let inner = js_sys::Reflect::construct(&constructor, &args)?; - let ready = js_sys::Reflect::get(&inner, &JsValue::from_str("ready"))? - .dyn_into::() - .map_err(|_| js_error("WebTransport.ready is not a Promise"))?; - JsFuture::from(ready) - .await - .map_err(|error| js_error(format!("WebTransport ready failed: {error:?}")))?; - Ok(Self { - inner, - incoming_reader: Rc::new(RefCell::new(None)), - }) - } - - async fn open_uni(&self) -> Result { - let create_stream = js_sys::Reflect::get( - &self.inner, - &JsValue::from_str("createUnidirectionalStream"), - )? - .dyn_into::() - .map_err(|_| js_error("createUnidirectionalStream not a function"))?; - let stream_promise = create_stream - .call0(&self.inner)? - .dyn_into::() - .map_err(|_| js_error("createUnidirectionalStream did not return a Promise"))?; - let stream = JsFuture::from(stream_promise).await?; - let writable = resolve_stream_writable(&stream)?; - let writer = js_sys::Reflect::get(&writable, &JsValue::from_str("getWriter")) - .map_err(|_| js_error("missing getWriter"))? - .dyn_into::() - .map_err(|_| js_error("getWriter not a function"))? - .call0(&writable) - .map_err(|_| js_error("getWriter call failed"))?; - Ok(BrowserSendStream { writer }) - } - - async fn accept_uni(&self) -> Result, JsValue> { - let streams_reader = if let Some(reader) = self.incoming_reader.borrow().clone() { - reader - } else { - let incoming = js_sys::Reflect::get( - &self.inner, - &JsValue::from_str("incomingUnidirectionalStreams"), - )?; - let reader = js_sys::Reflect::get(&incoming, &JsValue::from_str("getReader")) - .map_err(|_| js_error("missing getReader"))? - .dyn_into::() - .map_err(|_| js_error("getReader not a function"))? - .call0(&incoming) - .map_err(|_| js_error("getReader call failed"))?; - *self.incoming_reader.borrow_mut() = Some(reader.clone()); - reader - }; - - let read = js_sys::Reflect::get(&streams_reader, &JsValue::from_str("read")) - .map_err(|_| js_error("missing read"))? - .dyn_into::() - .map_err(|_| js_error("read not a function"))?; - let promise = read - .call0(&streams_reader) - .map_err(|_| js_error("read call failed"))? - .unchecked_into::(); - let result = JsFuture::from(promise).await.map_err(|error| { - log_stream_error_code(&error, "accept_uni"); - js_error(format!("accept stream failed: {error:?}")) - })?; - if js_sys::Reflect::get(&result, &JsValue::from_str("done")) - .ok() - .and_then(|value| value.as_bool()) - .unwrap_or(false) - { - return Ok(None); - } - - let stream = js_sys::Reflect::get(&result, &JsValue::from_str("value")) - .map_err(|_| js_error("missing value"))?; - let readable = resolve_stream_readable(&stream)?; - let reader = js_sys::Reflect::get(&readable, &JsValue::from_str("getReader")) - .map_err(|_| js_error("missing stream getReader"))? - .dyn_into::() - .map_err(|_| js_error("stream getReader not a function"))? - .call0(&readable) - .map_err(|_| js_error("stream getReader call failed"))?; - Ok(Some(BrowserRecvStream { reader })) - } - - fn close(&self) { - if let Some(reader) = self.incoming_reader.borrow_mut().take() { - release_reader_lock(&reader); - } - if let Ok(close) = js_sys::Reflect::get(&self.inner, &JsValue::from_str("close")) - .and_then(|value| value.dyn_into::()) - { - let _ = close.call1(&self.inner, &js_sys::Object::new()); - } - } -} - -impl BrowserSendStream { - pub(crate) async fn write_all(&mut self, bytes: &[u8]) -> Result<(), JsValue> { - let write = js_sys::Reflect::get(&self.writer, &JsValue::from_str("write")) - .map_err(|_| js_error("missing write"))? - .dyn_into::() - .map_err(|_| js_error("write not a function"))?; - let promise = write - .call1(&self.writer, &js_sys::Uint8Array::from(bytes)) - .map_err(|error| js_error(format!("write failed: {error:?}")))? - .unchecked_into::(); - JsFuture::from(promise).await.map(|_| ()) - } - - pub(crate) async fn finish(&mut self) -> Result<(), JsValue> { - let close = js_sys::Reflect::get(&self.writer, &JsValue::from_str("close")) - .map_err(|_| js_error("missing close"))? - .dyn_into::() - .map_err(|_| js_error("close not a function"))?; - let promise = close - .call0(&self.writer) - .map_err(|error| js_error(format!("close failed: {error:?}")))? - .unchecked_into::(); - JsFuture::from(promise).await.map(|_| ()) - } - - pub(crate) fn reset(&mut self, code: u32) -> Result<(), JsValue> { - let abort = js_sys::Reflect::get(&self.writer, &JsValue::from_str("abort")) - .map_err(|_| js_error("missing abort"))? - .dyn_into::() - .map_err(|_| js_error("abort not a function"))?; - let _ = abort.call1(&self.writer, &JsValue::from_f64(code as f64))?; - Ok(()) - } - - pub(crate) fn release(&self) { - release_writer_lock(&self.writer); - } -} - -impl BrowserRecvStream { - pub(crate) async fn read_chunk(&mut self) -> Result>, JsValue> { - let read = js_sys::Reflect::get(&self.reader, &JsValue::from_str("read")) - .map_err(|_| js_error("missing read"))? - .dyn_into::() - .map_err(|_| js_error("read not a function"))?; - let promise = read - .call0(&self.reader) - .map_err(|_| js_error("read call failed"))? - .unchecked_into::(); - let result = JsFuture::from(promise).await?; - if js_sys::Reflect::get(&result, &JsValue::from_str("done")) - .ok() - .and_then(|value| value.as_bool()) - .unwrap_or(true) - { - return Ok(None); - } - let value = js_sys::Reflect::get(&result, &JsValue::from_str("value")) - .map_err(|_| js_error("missing value"))?; - Ok(Some(js_sys::Uint8Array::new(&value).to_vec())) - } - - #[allow(dead_code)] - pub(crate) fn stop(self, code: u32) -> Result<(), JsValue> { - let cancel = js_sys::Reflect::get(&self.reader, &JsValue::from_str("cancel")) - .map_err(|_| js_error("missing cancel"))? - .dyn_into::() - .map_err(|_| js_error("cancel not a function"))?; - let _ = cancel.call1(&self.reader, &JsValue::from_f64(code as f64))?; - Ok(()) - } - - pub(crate) fn release(&self) { - release_reader_lock(&self.reader); - } -} - /// Releases a writer's lock so an abandoned writer isn't treated as an abort (which sends STOP_SENDING). pub(crate) fn release_writer_lock(writer: &JsValue) { if let Ok(release) = js_sys::Reflect::get(writer, &JsValue::from_str("releaseLock")) @@ -340,10 +118,12 @@ enum FrameOutcome { */ #[derive(Clone)] pub struct WasmTransport { - connection: BrowserConnection, + inner: JsValue, max_message_size: u32, - /// Current incoming unidirectional stream, shared across handshake and receive loops. - stream_reader: Rc>>, + /// Reader over `incoming_unidirectional_streams()` (a singleton stream of streams). + streams_reader: Rc>>, + /// Reader over the host's current uni-directional stream, if one is open. + stream_reader: Rc>>, /// Bytes already read from the current stream but not yet consumed as a frame. buffer: Rc>>, /// Set to `true` when `open_next_stream` succeeds; cleared after the first frame is parsed. @@ -369,14 +149,61 @@ impl WasmTransport { max_message_size: u32, configured_limits: Option, ) -> Result { - let connection = BrowserConnection::connect(url, cert_hashes).await?; + let ctor = js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("WebTransport"))? + .dyn_into::() + .map_err(|_| js_error("WebTransport not available"))?; + let args = js_sys::Array::new(); + args.push(&JsValue::from_str(url)); + + if let Some(hashes) = cert_hashes { + let wt_hashes = js_sys::Array::new(); + for h in hashes { + let (algo, hex_val) = match h.split_once(':') { + Some((algo, hex_val)) => (algo, hex_val), + None => ("sha-256", h.as_str()), + }; + + if let Ok(bytes) = hex::decode(hex_val) { + let hash = js_sys::Object::new(); + js_sys::Reflect::set( + &hash, + &JsValue::from_str("algorithm"), + &JsValue::from_str(algo), + )?; + js_sys::Reflect::set( + &hash, + &JsValue::from_str("value"), + &js_sys::Uint8Array::from(&bytes[..]), + )?; + wt_hashes.push(&hash); + } + } + if wt_hashes.length() > 0 { + let opts = js_sys::Object::new(); + js_sys::Reflect::set( + &opts, + &JsValue::from_str("serverCertificateHashes"), + &wt_hashes, + )?; + args.push(&opts); + } + }; + + let transport = js_sys::Reflect::construct(&ctor, &args)?; + let ready = js_sys::Reflect::get(&transport, &JsValue::from_str("ready"))? + .dyn_into::() + .map_err(|_| js_error("WebTransport.ready is not a Promise"))?; + JsFuture::from(ready) + .await + .map_err(|e| js_error(format!("WebTransport ready failed: {:?}", e)))?; let transport_limits = DecodeLimits::for_transport_message_size(max_message_size as u64); let decode_limits = configured_limits .map(|limits| restrict_decode_limits(limits, transport_limits)) .unwrap_or(transport_limits); Ok(Self { - connection, + inner: transport, max_message_size, + streams_reader: Rc::new(RefCell::new(None)), stream_reader: Rc::new(RefCell::new(None)), buffer: Rc::new(RefCell::new(Vec::new())), new_stream_frame: Rc::new(Cell::new(false)), @@ -387,7 +214,7 @@ impl WasmTransport { } pub fn inner(&self) -> &JsValue { - &self.connection.inner + &self.inner } pub fn set_type_map(&self, type_map: &TypeMap) { @@ -423,55 +250,154 @@ impl WasmTransport { // in accept_uni() until the authentication deadline. The bytes are // already the canonical MTP self-framed value, so no extra stream // length prefix is added here. - let mut stream = self.connection.open_uni().await?; - if let Err(e) = stream.write_all(frame).await { + let create_stream = js_sys::Reflect::get( + &self.inner, + &JsValue::from_str("createUnidirectionalStream"), + )? + .dyn_into::() + .map_err(|_| js_error("createUnidirectionalStream not a function"))?; + let stream_promise = create_stream + .call0(&self.inner)? + .dyn_into::() + .map_err(|_| js_error("createUnidirectionalStream did not return a Promise"))?; + let stream = JsFuture::from(stream_promise).await?; + let writable_or_stream = resolve_stream_writable(&stream)?; + let writer_val = js_sys::Reflect::get(&writable_or_stream, &JsValue::from_str("getWriter")) + .map_err(|_| js_error("missing getWriter"))? + .dyn_into::() + .map_err(|_| js_error("getWriter not a function"))? + .call0(&writable_or_stream) + .map_err(|_| js_error("getWriter call failed"))?; + + let chunk = js_sys::Uint8Array::from(frame); + + let write_fn = js_sys::Reflect::get(&writer_val, &JsValue::from_str("write")) + .map_err(|_| js_error("missing write"))? + .dyn_into::() + .map_err(|_| js_error("write not a function"))?; + let write_promise = write_fn + .call1(&writer_val, &chunk) + .map_err(|e| js_error(format!("write failed: {:?}", e)))?; + if let Err(e) = JsFuture::from(write_promise.unchecked_into::()).await { log_stream_error_code(&e, "send_frame write"); - stream.release(); + release_writer_lock(&writer_val); return Err(e); } - if let Err(e) = stream.finish().await { + let close_fn = js_sys::Reflect::get(&writer_val, &JsValue::from_str("close")) + .map_err(|_| js_error("missing close"))? + .dyn_into::() + .map_err(|_| js_error("close not a function"))?; + let close_promise = close_fn + .call0(&writer_val) + .map_err(|e| js_error(format!("close failed: {:?}", e)))?; + if let Err(e) = JsFuture::from(close_promise.unchecked_into::()).await { // The frame was already written; do not retry it merely because // FIN failed, as that would duplicate the MTP frame. log_stream_error_code(&e, "send_frame close"); } - stream.release(); + release_writer_lock(&writer_val); Ok(()) } + /// Get (creating once) the reader over `incoming_unidirectional_streams()`. + fn ensure_streams_reader(&self) -> Result { + if let Some(reader) = self.streams_reader.borrow().clone() { + return Ok(reader); + } + let incoming = js_sys::Reflect::get( + &self.inner, + &JsValue::from_str("incomingUnidirectionalStreams"), + )?; + let reader = js_sys::Reflect::get(&incoming, &JsValue::from_str("getReader")) + .map_err(|_| js_error("missing getReader"))? + .dyn_into::() + .map_err(|_| js_error("getReader not a function"))? + .call0(&incoming) + .map_err(|_| js_error("getReader call failed"))?; + *self.streams_reader.borrow_mut() = Some(reader.clone()); + Ok(reader) + } + /// Accept the next incoming uni-directional stream and make it current. /// Returns `false` if the incoming-streams readable has ended. async fn open_next_stream(&self) -> Result { - let Some(stream) = self.connection.accept_uni().await? else { - return Ok(false); + let streams_reader = self.ensure_streams_reader()?; + + let read_fn = js_sys::Reflect::get(&streams_reader, &JsValue::from_str("read")) + .map_err(|_| js_error("missing read"))? + .dyn_into::() + .map_err(|_| js_error("read not a function"))?; + let promise = read_fn + .call0(&streams_reader) + .map_err(|_| js_error("read call failed"))? + .unchecked_into::(); + let result = match JsFuture::from(promise).await { + Ok(r) => r, + Err(e) => { + log_stream_error_code(&e, "open_next_stream accept"); + return Err(js_error(format!("accept stream failed: {:?}", e))); + } }; - *self.stream_reader.borrow_mut() = Some(stream); + let done = js_sys::Reflect::get(&result, &JsValue::from_str("done")) + .ok() + .and_then(|v| v.as_bool()) + .unwrap_or(false); + if done { + return Ok(false); + } + + let recv_stream = js_sys::Reflect::get(&result, &JsValue::from_str("value")) + .map_err(|_| js_error("missing value"))?; + let readable = resolve_stream_readable(&recv_stream)?; + let reader = js_sys::Reflect::get(&readable, &JsValue::from_str("getReader")) + .map_err(|_| js_error("missing stream getReader"))? + .dyn_into::() + .map_err(|_| js_error("stream getReader not a function"))? + .call0(&readable) + .map_err(|_| js_error("stream getReader call failed"))?; + + *self.stream_reader.borrow_mut() = Some(reader); self.new_stream_frame.set(true); Ok(true) } /// Read one chunk from the current stream. `Ok(None)` means the stream ended. async fn read_chunk(&self) -> Result>, JsValue> { - let mut stream = match self.stream_reader.borrow_mut().take() { - Some(stream) => stream, + let reader = match self.stream_reader.borrow().clone() { + Some(r) => r, None => return Ok(None), }; - let result = match stream.read_chunk().await { - Ok(result) => result, + + let read_fn = js_sys::Reflect::get(&reader, &JsValue::from_str("read")) + .map_err(|_| js_error("missing read"))? + .dyn_into::() + .map_err(|_| js_error("read not a function"))?; + let promise = read_fn + .call0(&reader) + .map_err(|_| js_error("read call failed"))? + .unchecked_into::(); + let result = match JsFuture::from(promise).await { + Ok(r) => r, Err(e) => { log_stream_error_code(&e, "read_chunk"); - stream.release(); return Err(js_error(format!("read failed: {:?}", e))); } }; - if result.is_some() { - *self.stream_reader.borrow_mut() = Some(stream); - } else { - stream.release(); + + let done = js_sys::Reflect::get(&result, &JsValue::from_str("done")) + .ok() + .and_then(|v| v.as_bool()) + .unwrap_or(true); + if done { + return Ok(None); } - Ok(result) + + let value = js_sys::Reflect::get(&result, &JsValue::from_str("value")) + .map_err(|_| js_error("missing value"))?; + Ok(Some(js_sys::Uint8Array::new(&value).to_vec())) } /// Try to pull one complete frame out of the buffer without reading more. @@ -533,7 +459,9 @@ impl WasmTransport { } None => { // Stream finished; release the reader's lock to avoid a spurious cancel. - // `read_chunk` releases the raw stream lock on clean FIN. + if let Some(reader) = self.stream_reader.borrow_mut().take() { + release_reader_lock(&reader); + } // A frame is never allowed to span stream boundaries. The // native persistent-stream sender packs frames on one // stream, while the WASM sender uses one stream per frame; @@ -592,17 +520,15 @@ impl WasmTransport { /// Pipe-aware receive loop. Identical to `receive_loop` but detects /// `PipeRequest` as the first frame on a new incoming stream and routes /// the stream to `on_pipe` instead of `on_message`. - pub async fn receive_loop_with_pipes( + pub async fn receive_loop_with_pipes( &self, mut on_message: F, mut on_error: H, mut on_pipe: G, - mut pipe_is_expected: I, ) where F: FnMut(JsValue), G: FnMut(crate::pipe::PipeReader), H: FnMut(JsValue), - I: FnMut(u32) -> bool, { loop { match self.next_frame(self.max_message_size).await { @@ -624,44 +550,36 @@ impl WasmTransport { if is_first { self.new_stream_frame.set(false); - if let Some(comm) = comm.as_ref() { - let is_pipe_request = Some(comm.get_type()) == pipe_request_type; - let pipe_id = comm.id().filter(|id| *id != 0); - let is_expected = - is_pipe_request && pipe_id.is_some_and(&mut pipe_is_expected); - let disposition = - match classify_first_frame(is_pipe_request, comm.id(), is_expected) - { - Ok(disposition) => disposition, - Err(error) => { - on_error(JsValue::from_str(&error.to_string())); - self.close(); - break; - } - }; + if let Some(comm) = comm.as_ref() + && Some(comm.get_type()) == pipe_request_type + { + let Some(pipe_id) = comm.id().filter(|id| *id != 0) else { + on_error(JsValue::from_str( + "PipeRequest frame must contain a non-zero id", + )); + self.close(); + break; + }; + let description = comm + .get_str(mtp_codec::DataType::Description) + .unwrap_or("") + .to_string(); - if let FirstFrameDisposition::Pipe(pipe_id) = disposition { - let description = comm - .get_str(mtp_codec::DataType::Description) - .unwrap_or("") - .to_string(); + let pending = { + let mut buf = self.buffer.borrow_mut(); + std::mem::take(&mut *buf) + }; - let pending = { - let mut buf = self.buffer.borrow_mut(); - std::mem::take(&mut *buf) - }; - - if let Some(reader) = self.stream_reader.borrow_mut().take() { - let pipe_reader = crate::pipe::PipeReader::new( - reader, - pipe_id, - description, - pending, - ); - on_pipe(pipe_reader); - } - continue; + if let Some(reader) = self.stream_reader.borrow_mut().take() { + let pipe_reader = crate::pipe::PipeReader::new( + reader, + pipe_id, + description, + pending, + ); + on_pipe(pipe_reader); } + continue; } } @@ -717,7 +635,25 @@ impl WasmTransport { description: &str, ) -> Result { let _send_guard = self.send_lock.lock().await; - let mut stream = self.connection.open_uni().await?; + let create_stream = js_sys::Reflect::get( + &self.inner, + &JsValue::from_str("createUnidirectionalStream"), + )? + .dyn_into::() + .map_err(|_| js_error("createUnidirectionalStream not a function"))?; + let stream_promise = create_stream + .call0(&self.inner)? + .dyn_into::() + .map_err(|_| js_error("createUnidirectionalStream did not return a Promise"))?; + let stream = JsFuture::from(stream_promise).await?; + + let writable_or_stream = resolve_stream_writable(&stream)?; + let writer_val = js_sys::Reflect::get(&writable_or_stream, &JsValue::from_str("getWriter")) + .map_err(|_| js_error("missing getWriter"))? + .dyn_into::() + .map_err(|_| js_error("getWriter not a function"))? + .call0(&writable_or_stream) + .map_err(|_| js_error("getWriter call failed"))?; let type_map = self.type_map(); let request = mtp_codec::CommunicationValue::new_with_type_map( @@ -733,21 +669,37 @@ impl WasmTransport { .to_bytes() .map_err(|e| js_error(format!("encode failed: {}", e)))?; - if let Err(e) = stream.write_all(&frame_bytes).await { + let chunk = js_sys::Uint8Array::from(&frame_bytes[..]); + let write_fn = js_sys::Reflect::get(&writer_val, &JsValue::from_str("write")) + .map_err(|_| js_error("missing write"))? + .dyn_into::() + .map_err(|_| js_error("write not a function"))?; + let write_promise = write_fn + .call1(&writer_val, &chunk) + .map_err(|e| js_error(format!("write failed: {:?}", e)))?; + if let Err(e) = JsFuture::from(write_promise.unchecked_into::()).await { log_stream_error_code(&e, "open_pipe write"); - stream.release(); + release_writer_lock(&writer_val); return Err(e); } - Ok(crate::pipe::PipeWriter::new(stream, pipe_id)) + Ok(crate::pipe::PipeWriter::new(writer_val, pipe_id)) } pub fn close(&self) { // Release reader locks before closing so they aren't treated as cancels. if let Some(reader) = self.stream_reader.borrow_mut().take() { - reader.release(); + release_reader_lock(&reader); + } + if let Some(reader) = self.streams_reader.borrow_mut().take() { + release_reader_lock(&reader); + } + + if let Ok(close) = js_sys::Reflect::get(&self.inner, &JsValue::from_str("close")) + .and_then(|value| value.dyn_into::()) + { + let _ = close.call1(&self.inner, &js_sys::Object::new()); } - self.connection.close(); } }