Compare commits

...
Author SHA1 Message Date
3e4447cabd Update dependency jscpd to v5.0.15
Some checks failed
renovate/artifacts Artifact file update failure
renovate/stability-days Updates have met minimum release age requirement
CI / checks (pull_request) Failing after 2s
2026-08-27 17:01:05 +03:00
e83cd132a2
feat(wasm, native, h3): make wasm, native and h3 use unified interface
Some checks failed
CI / checks (push) Failing after 2s
2026-08-27 15:31:55 +02:00
14 changed files with 739 additions and 400 deletions

View file

@ -78,9 +78,41 @@ pub struct PipeRequest {
pub(crate) pipe_id: u32, pub(crate) pipe_id: u32,
pub(crate) description: String, pub(crate) description: String,
pub(crate) sender: Sender, pub(crate) sender: Sender,
pub(crate) receiver: Receiver,
pub(crate) dispatcher: Arc<PipeDispatcher>, pub(crate) dispatcher: Arc<PipeDispatcher>,
} }
#[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")] #[cfg(feature = "pipes")]
impl PipeRequest { impl PipeRequest {
pub fn id(&self) -> u32 { pub fn id(&self) -> u32 {
@ -92,6 +124,10 @@ impl PipeRequest {
} }
pub async fn accept(self) -> Result<mtp_transport::PipeReader, PipeError> { pub async fn accept(self) -> Result<mtp_transport::PipeReader, 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(); let (pipe_tx, pipe_rx) = tokio::sync::oneshot::channel();
{ {
let mut pending = self.dispatcher.pending_pipes.lock().await; let mut pending = self.dispatcher.pending_pipes.lock().await;
@ -115,7 +151,10 @@ impl PipeRequest {
let timeout = self.dispatcher.policy.read_timeout; let timeout = self.dispatcher.policy.read_timeout;
match tokio::time::timeout(timeout, pipe_rx).await { match tokio::time::timeout(timeout, pipe_rx).await {
Ok(Ok(reader)) => Ok(reader), Ok(Ok(reader)) => {
expected_pipe.disarm();
Ok(reader)
}
Ok(Err(_)) => { Ok(Err(_)) => {
self.dispatcher self.dispatcher
.pending_pipes .pending_pipes
@ -413,6 +452,7 @@ pub(crate) async fn run_dispatcher(
pipe_id, pipe_id,
description, description,
sender: sender.clone(), sender: sender.clone(),
receiver: receiver.clone(),
dispatcher: dispatcher.clone(), dispatcher: dispatcher.clone(),
}; };
let _ = pipe_req_tx.send(req).await; let _ = pipe_req_tx.send(req).await;

View file

@ -164,6 +164,9 @@ pub enum CommunicationError {
#[error("Stream Error")] #[error("Stream Error")]
StreamError, StreamError,
#[error("Stream failed after delivery may have started")]
DeliveryUnknown,
#[error("Stream Error: {0}")] #[error("Stream Error: {0}")]
#[cfg(not(target_arch = "wasm32"))] #[cfg(not(target_arch = "wasm32"))]
StreamWriteError(#[from] wtransport::error::StreamWriteError), StreamWriteError(#[from] wtransport::error::StreamWriteError),
@ -182,6 +185,38 @@ pub enum CommunicationError {
Other(String), 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<u32>,
pipe_is_expected: bool,
) -> Result<FirstFrameDisposition, CommunicationError> {
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) ---- // ---- manual PartialEq (quinn / wtransport types don't impl PartialEq) ----
impl PartialEq for CommunicationError { impl PartialEq for CommunicationError {
@ -212,6 +247,7 @@ impl PartialEq for CommunicationError {
(Self::ReadExactError(_), Self::ReadExactError(_)) => true, (Self::ReadExactError(_), Self::ReadExactError(_)) => true,
(Self::StreamClosed, Self::StreamClosed) => true, (Self::StreamClosed, Self::StreamClosed) => true,
(Self::StreamError, Self::StreamError) => true, (Self::StreamError, Self::StreamError) => true,
(Self::DeliveryUnknown, Self::DeliveryUnknown) => true,
#[cfg(not(target_arch = "wasm32"))] #[cfg(not(target_arch = "wasm32"))]
(Self::StreamWriteError(_), Self::StreamWriteError(_)) => true, (Self::StreamWriteError(_), Self::StreamWriteError(_)) => true,
#[cfg(not(target_arch = "wasm32"))] #[cfg(not(target_arch = "wasm32"))]

View file

@ -73,7 +73,7 @@ pub struct MTPConnection<
#[cfg(feature = "pipes")] #[cfg(feature = "pipes")]
pub(crate) app_rx: Mutex<mpsc::Receiver<Result<CommunicationValue, CommunicationError>>>, pub(crate) app_rx: Mutex<mpsc::Receiver<Result<CommunicationValue, CommunicationError>>>,
#[cfg(feature = "pipes")] #[cfg(feature = "pipes")]
pub(crate) pipe_req_rx: Mutex<mpsc::Receiver<PipeRequest<S, P>>>, pub(crate) pipe_req_rx: Mutex<mpsc::Receiver<PipeRequest<S, R, P>>>,
#[cfg(feature = "pipes")] #[cfg(feature = "pipes")]
pub(crate) pipe_dispatcher: Arc<PipeDispatcher<P>>, pub(crate) pipe_dispatcher: Arc<PipeDispatcher<P>>,
#[cfg(not(feature = "pipes"))] #[cfg(not(feature = "pipes"))]
@ -381,7 +381,7 @@ where
}) })
} }
pub async fn receive_pipe(&self) -> Result<PipeRequest<S, P>, CommunicationError> { pub async fn receive_pipe(&self) -> Result<PipeRequest<S, R, P>, CommunicationError> {
self.pipe_req_rx self.pipe_req_rx
.lock() .lock()
.await .await

View file

@ -27,6 +27,10 @@ pub trait PipeReceiver<P>: Clone + Send + Sync + 'static
where where
P: tokio::io::AsyncRead + Send + Unpin + 'static, 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( fn receive_pipe_event(
&self, &self,
) -> impl std::future::Future<Output = Result<TransportEvent<P>, CommunicationError>> + Send; ) -> impl std::future::Future<Output = Result<TransportEvent<P>, CommunicationError>> + Send;
@ -52,6 +56,14 @@ impl PipeSender for mtp_transport::Sender {
} }
impl PipeReceiver<wtransport::RecvStream> for mtp_transport::Receiver { impl PipeReceiver<wtransport::RecvStream> 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( async fn receive_pipe_event(
&self, &self,
) -> Result<TransportEvent<wtransport::RecvStream>, CommunicationError> { ) -> Result<TransportEvent<wtransport::RecvStream>, CommunicationError> {
@ -87,6 +99,14 @@ where
C: mtp_transport::TransportConnection, C: mtp_transport::TransportConnection,
C::RecvStream: tokio::io::AsyncRead + Send + Unpin + 'static, 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( async fn receive_pipe_event(
&self, &self,
) -> Result<TransportEvent<C::RecvStream>, CommunicationError> { ) -> Result<TransportEvent<C::RecvStream>, CommunicationError> {
@ -152,16 +172,60 @@ where
} }
} }
pub struct PipeRequest<S, P> { pub struct PipeRequest<S, R, P> {
pub(crate) pipe_id: u32, pub(crate) pipe_id: u32,
pub(crate) description: String, pub(crate) description: String,
pub(crate) sender: S, pub(crate) sender: S,
pub(crate) receiver: R,
pub(crate) dispatcher: Arc<PipeDispatcher<P>>, pub(crate) dispatcher: Arc<PipeDispatcher<P>>,
} }
impl<S, P> PipeRequest<S, P> struct ExpectedPipeGuard<R, P>
where
R: PipeReceiver<P>,
P: tokio::io::AsyncRead + Send + Unpin + 'static,
{
receiver: R,
pipe_id: u32,
armed: bool,
_stream: std::marker::PhantomData<P>,
}
impl<R, P> ExpectedPipeGuard<R, P>
where
R: PipeReceiver<P>,
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<R, P> Drop for ExpectedPipeGuard<R, P>
where
R: PipeReceiver<P>,
P: tokio::io::AsyncRead + Send + Unpin + 'static,
{
fn drop(&mut self) {
if self.armed {
self.receiver.cancel_expected_pipe(self.pipe_id);
}
}
}
impl<S, R, P> PipeRequest<S, R, P>
where where
S: PipeSender, S: PipeSender,
R: PipeReceiver<P>,
P: tokio::io::AsyncRead + Send + Unpin + 'static, P: tokio::io::AsyncRead + Send + Unpin + 'static,
{ {
pub fn id(&self) -> u32 { pub fn id(&self) -> u32 {
@ -173,6 +237,10 @@ where
} }
pub async fn accept(self) -> Result<PipeReader<P>, PipeError> { pub async fn accept(self) -> Result<PipeReader<P>, PipeError> {
self.receiver
.expect_pipe(self.pipe_id)
.map_err(PipeError::from)?;
let mut expected_pipe = ExpectedPipeGuard::<R, P>::new(self.receiver.clone(), self.pipe_id);
let (pipe_tx, pipe_rx) = tokio::sync::oneshot::channel(); let (pipe_tx, pipe_rx) = tokio::sync::oneshot::channel();
self.dispatcher self.dispatcher
.pending_pipes .pending_pipes
@ -196,7 +264,10 @@ where
} }
match tokio::time::timeout(self.dispatcher.policy.read_timeout, pipe_rx).await { match tokio::time::timeout(self.dispatcher.policy.read_timeout, pipe_rx).await {
Ok(Ok(reader)) => Ok(reader), Ok(Ok(reader)) => {
expected_pipe.disarm();
Ok(reader)
}
Ok(Err(_)) => { Ok(Err(_)) => {
self.dispatcher self.dispatcher
.pending_pipes .pending_pipes
@ -361,7 +432,7 @@ pub(crate) async fn run_dispatcher<S, R, P>(
receiver: R, receiver: R,
sender: S, sender: S,
app_tx: mpsc::Sender<Result<CommunicationValue, CommunicationError>>, app_tx: mpsc::Sender<Result<CommunicationValue, CommunicationError>>,
pipe_req_tx: mpsc::Sender<PipeRequest<S, P>>, pipe_req_tx: mpsc::Sender<PipeRequest<S, R, P>>,
dispatcher: Arc<PipeDispatcher<P>>, dispatcher: Arc<PipeDispatcher<P>>,
) where ) where
S: PipeSender, S: PipeSender,
@ -388,6 +459,7 @@ pub(crate) async fn run_dispatcher<S, R, P>(
.unwrap_or("") .unwrap_or("")
.to_owned(), .to_owned(),
sender: sender.clone(), sender: sender.clone(),
receiver: receiver.clone(),
dispatcher: dispatcher.clone(), dispatcher: dispatcher.clone(),
}; };
let _ = pipe_req_tx.send(request).await; let _ = pipe_req_tx.send(request).await;
@ -446,6 +518,7 @@ pub(crate) async fn run_dispatcher<S, R, P>(
pipe_id, pipe_id,
description: reader.description().to_owned(), description: reader.description().to_owned(),
sender: sender.clone(), sender: sender.clone(),
receiver: receiver.clone(),
dispatcher: dispatcher.clone(), dispatcher: dispatcher.clone(),
}; };
let _ = pipe_req_tx.send(request).await; let _ = pipe_req_tx.send(request).await;

View file

@ -57,14 +57,14 @@ impl TransportSendStream for H3TransportSender {
self.stream self.stream
.write_all(buf) .write_all(buf)
.await .await
.map_err(|_| CommunicationError::StreamError)?; .map_err(|_| CommunicationError::DeliveryUnknown)?;
// Control/authentication frames use a persistent stream. h3 keeps // Control/authentication frames use a persistent stream. h3 keeps
// those writes buffered until flushed; without this the peer can wait // those writes buffered until flushed; without this the peer can wait
// for the challenge while the server waits for its proof. // for the challenge while the server waits for its proof.
self.stream self.stream
.flush() .flush()
.await .await
.map_err(|_| CommunicationError::StreamError) .map_err(|_| CommunicationError::DeliveryUnknown)
} }
async fn finish(&mut self) -> Result<(), CommunicationError> { async fn finish(&mut self) -> Result<(), CommunicationError> {
@ -73,6 +73,11 @@ impl TransportSendStream for H3TransportSender {
.await .await
.map_err(|_| CommunicationError::StreamError) .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] #[async_trait::async_trait]
@ -140,6 +145,11 @@ 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 { impl tokio::io::AsyncWrite for H3TransportSender {

View file

@ -72,7 +72,7 @@
}, },
"devDependencies": { "devDependencies": {
"@types/node": "^26.0.1", "@types/node": "^26.0.1",
"jscpd": "5.0.14", "jscpd": "5.0.15",
"typescript": "^7.0.0" "typescript": "^7.0.0"
}, },
"dependencies": { "dependencies": {

View file

@ -4,6 +4,10 @@ use crate::framing::RetryClassifier;
use crate::pipe::PipeReader; use crate::pipe::PipeReader;
use mtp_codec::{CommunicationValue, DecodeError, DecodeLimits, EncodeLimits, TypeMap}; use mtp_codec::{CommunicationValue, DecodeError, DecodeLimits, EncodeLimits, TypeMap};
use mtp_common::CommunicationError; 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::ops::Deref;
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::atomic::{AtomicU64, Ordering};
@ -288,15 +292,15 @@ impl Sender {
Ok(Ok(())) => Ok(()), Ok(Ok(())) => Ok(()),
Ok(Err(wtransport::error::StreamWriteError::Stopped(code))) => { Ok(Err(wtransport::error::StreamWriteError::Stopped(code))) => {
warn!("[Sender] write failed: peer sent STOP_SENDING (error code {code})"); warn!("[Sender] write failed: peer sent STOP_SENDING (error code {code})");
Err(CommunicationError::StreamClosed) Err(CommunicationError::DeliveryUnknown)
} }
Ok(Err(other)) => { Ok(Err(other)) => {
warn!("[Sender] write failed: {other}"); warn!("[Sender] write failed: {other}");
Err(CommunicationError::StreamError) Err(CommunicationError::DeliveryUnknown)
} }
Err(_) => { Err(_) => {
warn!("[Sender] write timed out (len={})", bytes.len()); warn!("[Sender] write timed out (len={})", bytes.len());
Err(CommunicationError::StreamError) Err(CommunicationError::DeliveryUnknown)
} }
} }
} }
@ -398,15 +402,15 @@ impl Sender {
Ok(Ok(())) => Ok(()), Ok(Ok(())) => Ok(()),
Ok(Err(wtransport::error::StreamWriteError::Stopped(code))) => { Ok(Err(wtransport::error::StreamWriteError::Stopped(code))) => {
warn!("[Sender] finish failed: peer sent STOP_SENDING (error code {code})"); warn!("[Sender] finish failed: peer sent STOP_SENDING (error code {code})");
Err(CommunicationError::StreamClosed) Err(CommunicationError::DeliveryUnknown)
} }
Ok(Err(other)) => { Ok(Err(other)) => {
warn!("[Sender] finish failed: {other}"); warn!("[Sender] finish failed: {other}");
Err(CommunicationError::StreamError) Err(CommunicationError::DeliveryUnknown)
} }
Err(_) => { Err(_) => {
warn!("[Sender] finish timed out"); warn!("[Sender] finish timed out");
Err(CommunicationError::StreamError) Err(CommunicationError::DeliveryUnknown)
} }
} }
} }
@ -745,6 +749,8 @@ struct ReceiverInner {
max_message_size: Arc<AtomicU64>, max_message_size: Arc<AtomicU64>,
type_map: Arc<RwLock<TypeMap>>, type_map: Arc<RwLock<TypeMap>>,
decode_rejections: Arc<DecodeRejectionCounters>, decode_rejections: Arc<DecodeRejectionCounters>,
#[cfg(feature = "pipes")]
expected_pipes: Arc<std::sync::Mutex<HashSet<u32>>>,
} }
impl Clone for Receiver { impl Clone for Receiver {
@ -831,6 +837,10 @@ impl Receiver {
let accept_type_map = type_map.clone(); let accept_type_map = type_map.clone();
let decode_rejections = Arc::new(DecodeRejectionCounters::default()); let decode_rejections = Arc::new(DecodeRejectionCounters::default());
let accept_decode_rejections = decode_rejections.clone(); 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 stream_limit = Arc::new(Semaphore::new(policy.max_concurrent_stream_tasks.max(1)));
let accept_stream_limit = stream_limit.clone(); let accept_stream_limit = stream_limit.clone();
debug!( debug!(
@ -900,6 +910,8 @@ impl Receiver {
let stream_max_message_size = accept_max_message_size.clone(); let stream_max_message_size = accept_max_message_size.clone();
let stream_type_map = accept_type_map.clone(); let stream_type_map = accept_type_map.clone();
let stream_decode_rejections = accept_decode_rejections.clone(); let stream_decode_rejections = accept_decode_rejections.clone();
#[cfg(feature = "pipes")]
let stream_expected_pipes = accept_expected_pipes.clone();
tokio::spawn(async move { tokio::spawn(async move {
let _permit = permit; let _permit = permit;
@ -935,38 +947,54 @@ impl Receiver {
#[cfg(feature = "pipes")] #[cfg(feature = "pipes")]
{ {
if msg.is_type(mtp_codec::CommunicationType::PipeRequest) if frame_count == 1 {
&& frame_count == 1 let is_pipe_request = msg.is_type(
{ mtp_codec::CommunicationType::PipeRequest,
let Some(pipe_id) = msg.id().filter(|id| *id != 0) else { );
let error = CommunicationError::Other( let pipe_id = msg.id().filter(|id| *id != 0);
"PipeRequest frame must contain a non-zero id".into(), let pipe_is_expected = is_pipe_request && pipe_id.is_some_and(|pipe_id| {
); stream_expected_pipes
let _ = msg_tx_stream.send(Err(error.clone())).await; .lock()
stream_handle.close(Some(error)); .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 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; 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 pipe_tx_stream
.send(pipe_reader)
.await
.is_err()
{
stream_handle.close(Some(
CommunicationError::StreamClosed,
));
} }
break;
} }
} }
@ -1111,6 +1139,8 @@ impl Receiver {
max_message_size, max_message_size,
type_map, type_map,
decode_rejections, decode_rejections,
#[cfg(feature = "pipes")]
expected_pipes,
}), }),
} }
} }
@ -1127,6 +1157,26 @@ impl Receiver {
*self.inner.type_map.write().await = type_map.clone(); *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. /// Return local counts for frames rejected by the structured decoder.
/// ///
/// These counters are intentionally local-only; peers continue to receive /// These counters are intentionally local-only; peers continue to receive

View file

@ -82,6 +82,10 @@ mod tests {
async fn finish(&mut self) -> Result<(), CommunicationError> { async fn finish(&mut self) -> Result<(), CommunicationError> {
Ok(()) Ok(())
} }
fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> {
Ok(())
}
} }
#[tokio::test] #[tokio::test]

View file

@ -10,8 +10,12 @@ use crate::{
framing::{RetryClassifier, write_frame}, framing::{RetryClassifier, write_frame},
}; };
use mtp_codec::{CommunicationValue, DataType, DecodeLimits, TypeMap}; use mtp_codec::{CommunicationValue, DataType, DecodeLimits, TypeMap};
use mtp_common::CommunicationError; use mtp_common::{CommunicationError, FirstFrameDisposition, classify_first_frame};
#[cfg(feature = "pipes")]
use std::collections::HashSet;
use std::sync::Arc; use std::sync::Arc;
#[cfg(feature = "pipes")]
use std::sync::Mutex as StdMutex;
use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::atomic::{AtomicU64, Ordering};
use tokio::sync::{Mutex, Notify, RwLock, Semaphore, mpsc}; use tokio::sync::{Mutex, Notify, RwLock, Semaphore, mpsc};
use tokio::time::{Instant, timeout, timeout_at}; use tokio::time::{Instant, timeout, timeout_at};
@ -85,9 +89,10 @@ impl<C: TransportConnection> GenericSender<C> {
) )
.await .await
.map_err(|_| CommunicationError::StreamError)??; .map_err(|_| CommunicationError::StreamError)??;
timeout(self.policy.write_timeout, stream.finish()) match timeout(self.policy.write_timeout, stream.finish()).await {
.await Ok(Ok(())) => Ok(()),
.map_err(|_| CommunicationError::StreamError)? Ok(Err(_)) | Err(_) => Err(CommunicationError::DeliveryUnknown),
}
} }
crate::SendMode::PersistentStream => { crate::SendMode::PersistentStream => {
let mut stream = self.persistent.lock().await; let mut stream = self.persistent.lock().await;
@ -204,6 +209,8 @@ pub struct GenericReceiver<C: TransportConnection> {
type_map: Arc<RwLock<TypeMap>>, type_map: Arc<RwLock<TypeMap>>,
queue_notify: Arc<Notify>, queue_notify: Arc<Notify>,
decode_rejections: Arc<DecodeRejectionCounters>, decode_rejections: Arc<DecodeRejectionCounters>,
#[cfg(feature = "pipes")]
expected_pipes: Arc<StdMutex<HashSet<u32>>>,
_accept_task: Arc<tokio::task::JoinHandle<()>>, _accept_task: Arc<tokio::task::JoinHandle<()>>,
} }
@ -219,6 +226,8 @@ impl<C: TransportConnection> Clone for GenericReceiver<C> {
type_map: self.type_map.clone(), type_map: self.type_map.clone(),
queue_notify: self.queue_notify.clone(), queue_notify: self.queue_notify.clone(),
decode_rejections: self.decode_rejections.clone(), decode_rejections: self.decode_rejections.clone(),
#[cfg(feature = "pipes")]
expected_pipes: self.expected_pipes.clone(),
_accept_task: self._accept_task.clone(), _accept_task: self._accept_task.clone(),
} }
} }
@ -254,6 +263,10 @@ impl<C: TransportConnection> GenericReceiver<C> {
let task_queue_notify = queue_notify.clone(); let task_queue_notify = queue_notify.clone();
let decode_rejections = Arc::new(DecodeRejectionCounters::default()); let decode_rejections = Arc::new(DecodeRejectionCounters::default());
let task_decode_rejections = decode_rejections.clone(); 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(); let task_accept_task_tx = tx.clone();
#[cfg(feature = "pipes")] #[cfg(feature = "pipes")]
let task_accept_task_pipe_tx = pipe_tx.clone(); let task_accept_task_pipe_tx = pipe_tx.clone();
@ -312,6 +325,8 @@ impl<C: TransportConnection> GenericReceiver<C> {
let connection = task_connection.clone(); let connection = task_connection.clone();
let type_map = task_type_map.clone(); let type_map = task_type_map.clone();
let decode_rejections = task_decode_rejections.clone(); let decode_rejections = task_decode_rejections.clone();
#[cfg(feature = "pipes")]
let expected_pipes = task_expected_pipes.clone();
tokio::spawn(async move { tokio::spawn(async move {
let _permit = permit; let _permit = permit;
let mut stream = stream; let mut stream = stream;
@ -443,37 +458,51 @@ impl<C: TransportConnection> GenericReceiver<C> {
#[cfg(feature = "pipes")] #[cfg(feature = "pipes")]
{ {
if message.is_type(mtp_codec::CommunicationType::PipeRequest) if frames == 1 {
&& frames == 1 let is_pipe_request =
{ message.is_type(mtp_codec::CommunicationType::PipeRequest);
let Some(pipe_id) = message.id().filter(|id| *id != 0) else { let pipe_id = message.id().filter(|id| *id != 0);
let error = CommunicationError::Other( let pipe_is_expected = is_pipe_request
"PipeRequest frame must contain a non-zero id".into(), && pipe_id.is_some_and(|pipe_id| {
); expected_pipes
let _ = tx.send(Err(error.clone())).await; .lock()
connection.close( .is_ok_and(|mut expected| expected.remove(&pipe_id))
policy.application_close_code, });
b"pipe request missing id", let disposition = match classify_first_frame(
); is_pipe_request,
break; message.id(),
}; pipe_is_expected,
let description = message ) {
.get_str(mtp_codec::DataType::Description) Ok(disposition) => disposition,
.unwrap_or("") Err(error) => {
.to_string(); let _ = tx.send(Err(error.clone())).await;
connection.close(
let pipe_reader = PipeReader { policy.application_close_code,
stream, b"pipe request missing id",
description, );
pipe_id, break;
}
}; };
tracing::debug!(pipe_id, description = %pipe_reader.description, "classified incoming pipe stream"); if let FirstFrameDisposition::Pipe(pipe_id) = disposition {
let description = message
.get_str(mtp_codec::DataType::Description)
.unwrap_or("")
.to_string();
if pipe_tx.send(pipe_reader).await.is_err() { let pipe_reader = PipeReader {
break; 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;
} }
return;
} }
} }
@ -517,6 +546,8 @@ impl<C: TransportConnection> GenericReceiver<C> {
type_map, type_map,
queue_notify, queue_notify,
decode_rejections, decode_rejections,
#[cfg(feature = "pipes")]
expected_pipes,
_accept_task: Arc::new(accept_task), _accept_task: Arc::new(accept_task),
} }
} }
@ -524,6 +555,25 @@ impl<C: TransportConnection> GenericReceiver<C> {
*self.ping_sender.write().await = Some(sender); *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. /// Switch from the handshake frame limit to the application frame limit.
pub fn set_max_message_size(&self, max_message_size: u64) { pub fn set_max_message_size(&self, max_message_size: u64) {
self.max_message_size self.max_message_size

View file

@ -17,6 +17,7 @@ use mtp_common::CommunicationError;
pub trait TransportSendStream: tokio::io::AsyncWrite + Send + Sync { pub trait TransportSendStream: tokio::io::AsyncWrite + Send + Sync {
async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError>; async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError>;
async fn finish(&mut self) -> 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. /// A readable unidirectional stream suitable for MTP frames.
@ -28,6 +29,9 @@ pub trait TransportSendStream: tokio::io::AsyncWrite + Send + Sync {
pub trait TransportRecvStream: tokio::io::AsyncRead + Send + Sync { pub trait TransportRecvStream: tokio::io::AsyncRead + Send + Sync {
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError>; async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError>;
async fn read_chunk(&mut self, max: usize) -> Result<Option<Vec<u8>>, CommunicationError>; async fn read_chunk(&mut self, max: usize) -> Result<Option<Vec<u8>>, CommunicationError>;
fn stop(self, code: u32) -> Result<(), CommunicationError>
where
Self: Sized;
} }
/// A QUIC/WebTransport connection that provides MTP's unidirectional streams. /// A QUIC/WebTransport connection that provides MTP's unidirectional streams.
@ -47,7 +51,7 @@ impl TransportSendStream for wtransport::SendStream {
async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError> { async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError> {
wtransport::SendStream::write_all(self, buf) wtransport::SendStream::write_all(self, buf)
.await .await
.map_err(|_| CommunicationError::StreamError) .map_err(|_| CommunicationError::DeliveryUnknown)
} }
async fn finish(&mut self) -> Result<(), CommunicationError> { async fn finish(&mut self) -> Result<(), CommunicationError> {
@ -55,14 +59,23 @@ impl TransportSendStream for wtransport::SendStream {
.await .await
.map_err(|_| CommunicationError::StreamError) .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] #[async_trait]
impl TransportRecvStream for wtransport::RecvStream { impl TransportRecvStream for wtransport::RecvStream {
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError> { async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError> {
wtransport::RecvStream::read_exact(self, buf) match wtransport::RecvStream::read_exact(self, buf).await {
.await Ok(()) => Ok(()),
.map_err(|_| CommunicationError::StreamError) Err(wtransport::error::StreamReadExactError::FinishedEarly(0)) => {
Err(CommunicationError::StreamClosed)
}
Err(_) => Err(CommunicationError::StreamError),
}
} }
async fn read_chunk(&mut self, max: usize) -> Result<Option<Vec<u8>>, CommunicationError> { async fn read_chunk(&mut self, max: usize) -> Result<Option<Vec<u8>>, CommunicationError> {
@ -76,6 +89,11 @@ impl TransportRecvStream for wtransport::RecvStream {
Err(_) => Err(CommunicationError::StreamError), Err(_) => Err(CommunicationError::StreamError),
} }
} }
fn stop(self, code: u32) -> Result<(), CommunicationError> {
wtransport::RecvStream::stop(self, wtransport::VarInt::from_u32(code));
Ok(())
}
} }
#[async_trait] #[async_trait]

View file

@ -4,7 +4,7 @@ use async_trait::async_trait;
use mtp_codec::CommunicationValue; use mtp_codec::CommunicationValue;
use mtp_common::CommunicationError; use mtp_common::CommunicationError;
use mtp_transport::{ use mtp_transport::{
GenericReceiver, GenericSender, Policy, TransportConnection, TransportEvent, GenericReceiver, GenericSender, Policy, SendMode, TransportConnection, TransportEvent,
TransportRecvStream, TransportSendStream, TransportRecvStream, TransportSendStream,
}; };
use std::sync::Arc; use std::sync::Arc;
@ -53,6 +53,10 @@ impl TransportSendStream for MockSendStream {
.await .await
.map_err(|_| CommunicationError::StreamError) .map_err(|_| CommunicationError::StreamError)
} }
fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> {
Ok(())
}
} }
struct MockRecvStream { struct MockRecvStream {
@ -89,6 +93,10 @@ impl TransportRecvStream for MockRecvStream {
Err(_) => Err(CommunicationError::StreamError), Err(_) => Err(CommunicationError::StreamError),
} }
} }
fn stop(self, _code: u32) -> Result<(), CommunicationError> {
Ok(())
}
} }
#[derive(Clone)] #[derive(Clone)]
@ -157,6 +165,7 @@ async fn test_open_pipe_and_receive_reader() -> Result<(), Box<dyn std::error::E
let sender = GenericSender::new(conn_a, policy.clone()); let sender = GenericSender::new(conn_a, policy.clone());
let receiver = GenericReceiver::new(conn_b, policy); let receiver = GenericReceiver::new(conn_b, policy);
receiver.expect_pipe(42)?;
let pipe_writer = sender.open_pipe(42, "test-pipe").await?; let pipe_writer = sender.open_pipe(42, "test-pipe").await?;
let pipe_reader = receiver.receive_pipe().await?; let pipe_reader = receiver.receive_pipe().await?;
@ -174,6 +183,7 @@ async fn test_pipe_raw_data_roundtrip() -> Result<(), Box<dyn std::error::Error>
let sender = GenericSender::new(conn_a, policy.clone()); let sender = GenericSender::new(conn_a, policy.clone());
let receiver = GenericReceiver::new(conn_b, policy); let receiver = GenericReceiver::new(conn_b, policy);
receiver.expect_pipe(1)?;
let mut pipe_writer = sender.open_pipe(1, "data-pipe").await?; let mut pipe_writer = sender.open_pipe(1, "data-pipe").await?;
let data = b"hello through the pipe"; let data = b"hello through the pipe";
@ -195,6 +205,7 @@ async fn test_pipe_large_payload() -> Result<(), Box<dyn std::error::Error>> {
let sender = GenericSender::new(conn_a, policy.clone()); let sender = GenericSender::new(conn_a, policy.clone());
let receiver = GenericReceiver::new(conn_b, policy); let receiver = GenericReceiver::new(conn_b, policy);
receiver.expect_pipe(7)?;
let mut pipe_writer = sender.open_pipe(7, "big-pipe").await?; let mut pipe_writer = sender.open_pipe(7, "big-pipe").await?;
let data: Vec<u8> = (0..256 * 1024).map(|i| (i % 256) as u8).collect(); let data: Vec<u8> = (0..256 * 1024).map(|i| (i % 256) as u8).collect();
@ -223,6 +234,7 @@ async fn test_receive_event_dispatches_pipe() -> Result<(), Box<dyn std::error::
let sender = GenericSender::new(conn_a, policy.clone()); let sender = GenericSender::new(conn_a, policy.clone());
let receiver = GenericReceiver::new(conn_b, policy); let receiver = GenericReceiver::new(conn_b, policy);
receiver.expect_pipe(99)?;
let mut pipe_writer = sender.open_pipe(99, "event-pipe").await?; let mut pipe_writer = sender.open_pipe(99, "event-pipe").await?;
match receiver.receive_event().await? { match receiver.receive_event().await? {
@ -261,22 +273,23 @@ async fn test_try_receive_pipe_returns_none_when_empty() -> Result<(), Box<dyn s
async fn test_regular_messages_still_work_alongside_pipes() -> Result<(), Box<dyn std::error::Error>> async fn test_regular_messages_still_work_alongside_pipes() -> Result<(), Box<dyn std::error::Error>>
{ {
let (conn_a, conn_b) = mock_connected_pair().await; let (conn_a, conn_b) = mock_connected_pair().await;
let policy = Arc::new(Policy::default()); let policy = Arc::new(Policy::default().with_send_mode(SendMode::SingleStreamPerMessage));
let sender = GenericSender::new(conn_a, policy.clone()); let sender = GenericSender::new(conn_a, policy.clone());
let receiver = GenericReceiver::new(conn_b, policy); let receiver = GenericReceiver::new(conn_b, policy);
let msg = CommunicationValue::new(mtp_codec::CommunicationType::BadRequest); let request = CommunicationValue::new(mtp_codec::CommunicationType::PipeRequest)
sender.send(&msg).await?; .with_id(1)
.add_typed_default(
let _pipe_writer = sender.open_pipe(1, "mixed-pipe").await?; mtp_codec::DataType::Description,
mtp_codec::DataValue::Str("mixed-pipe".into()),
);
sender.send(&request).await?;
let received = receiver.receive().await?; let received = receiver.receive().await?;
assert_eq!( assert!(received.is_type(mtp_codec::CommunicationType::PipeRequest));
received.get_type(),
mtp_codec::CommunicationType::BadRequest receiver.expect_pipe(1)?;
.try_to_id(&mtp_codec::TypeMap::latest()) let _pipe_writer = sender.open_pipe(1, "mixed-pipe").await?;
.unwrap()
);
let pipe_reader = receiver.receive_pipe().await?; let pipe_reader = receiver.receive_pipe().await?;
assert_eq!(pipe_reader.pipe_id(), 1); assert_eq!(pipe_reader.pipe_id(), 1);
@ -291,6 +304,8 @@ async fn test_multiple_pipes() -> Result<(), Box<dyn std::error::Error>> {
let sender = GenericSender::new(conn_a, policy.clone()); let sender = GenericSender::new(conn_a, policy.clone());
let receiver = GenericReceiver::new(conn_b, policy); 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 pw1 = sender.open_pipe(10, "first").await?;
let mut pw2 = sender.open_pipe(20, "second").await?; let mut pw2 = sender.open_pipe(20, "second").await?;

View file

@ -179,6 +179,7 @@ impl WasmClient {
let expired_pipe_creations = self.expired_pipe_creations.clone(); let expired_pipe_creations = self.expired_pipe_creations.clone();
let pending_pipes = self.pending_pipes.clone(); let pending_pipes = self.pending_pipes.clone();
let loop_pending_pipes = 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 on_pipe_request = self.on_pipe_request.clone();
let loop_pipe_creations = pending_pipe_creations.clone(); let loop_pipe_creations = pending_pipe_creations.clone();
let loop_expired_pipe_creations = expired_pipe_creations.clone(); let loop_expired_pipe_creations = expired_pipe_creations.clone();
@ -293,6 +294,12 @@ impl WasmClient {
let _ = entry.sender.send(Ok(pipe_reader)); 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; .await;
if connection_generation.get() != generation { if connection_generation.get() != generation {

View file

@ -1,9 +1,6 @@
use wasm_bindgen::JsCast;
use wasm_bindgen::prelude::*; use wasm_bindgen::prelude::*;
use wasm_bindgen_futures::JsFuture;
use crate::error::js_error; use crate::transport::{BrowserRecvStream, BrowserSendStream, log_stream_error_code};
use crate::transport::release_writer_lock;
#[wasm_bindgen(typescript_custom_section)] #[wasm_bindgen(typescript_custom_section)]
const PIPE_TS: &str = r#" const PIPE_TS: &str = r#"
@ -23,54 +20,41 @@ export interface PipeReader {
#[wasm_bindgen] #[wasm_bindgen]
pub struct PipeWriter { pub struct PipeWriter {
writer: JsValue, stream: BrowserSendStream,
pipe_id: u32, pipe_id: u32,
} }
impl PipeWriter { impl PipeWriter {
pub fn new(writer: JsValue, pipe_id: u32) -> Self { pub(crate) fn new(stream: BrowserSendStream, pipe_id: u32) -> Self {
Self { writer, pipe_id } Self { stream, pipe_id }
}
}
impl Drop for PipeWriter {
fn drop(&mut self) {
self.stream.release();
} }
} }
#[wasm_bindgen] #[wasm_bindgen]
impl PipeWriter { impl PipeWriter {
pub async fn write(&mut self, data: &[u8]) -> Result<(), JsValue> { pub async fn write(&mut self, data: &[u8]) -> Result<(), JsValue> {
let chunk = js_sys::Uint8Array::from(data); self.stream.write_all(data).await
let write_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("write"))
.map_err(|_| js_error("missing write"))?
.dyn_into::<js_sys::Function>()
.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::<js_sys::Promise>()).await?;
Ok(())
} }
pub async fn close(self) -> Result<(), JsValue> { pub async fn close(mut self) -> Result<(), JsValue> {
let close_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("close")) let result = self.stream.finish().await;
.map_err(|_| js_error("missing close"))? if let Err(error) = &result {
.dyn_into::<js_sys::Function>() log_stream_error_code(error, "pipe writer close");
.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::<js_sys::Promise>()).await {
crate::transport::log_stream_error_code(&e, "pipe writer close");
} }
release_writer_lock(&self.writer); self.stream.release();
Ok(()) result
} }
pub fn abort(&mut self) -> Result<(), JsValue> { pub fn abort(&mut self) -> Result<(), JsValue> {
let abort_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("abort")) let result = self.stream.reset(0);
.map_err(|_| js_error("missing abort"))? self.stream.release();
.dyn_into::<js_sys::Function>() result
.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 { pub fn pipe_id(&self) -> u32 {
@ -80,19 +64,26 @@ impl PipeWriter {
#[wasm_bindgen] #[wasm_bindgen]
pub struct PipeReader { pub struct PipeReader {
reader: JsValue, stream: BrowserRecvStream,
description: String, description: String,
pipe_id: u32, pipe_id: u32,
pending: Vec<u8>, pending: Vec<u8>,
finished: bool,
} }
impl PipeReader { impl PipeReader {
pub fn new(reader: JsValue, pipe_id: u32, description: String, pending: Vec<u8>) -> Self { pub(crate) fn new(
stream: BrowserRecvStream,
pipe_id: u32,
description: String,
pending: Vec<u8>,
) -> Self {
Self { Self {
reader, stream,
pipe_id, pipe_id,
description, description,
pending, pending,
finished: false,
} }
} }
} }
@ -105,27 +96,18 @@ impl PipeReader {
return Ok(js_sys::Uint8Array::from(&data[..]).into()); return Ok(js_sys::Uint8Array::from(&data[..]).into());
} }
let read_fn = js_sys::Reflect::get(&self.reader, &JsValue::from_str("read")) if self.finished {
.map_err(|_| js_error("missing read"))?
.dyn_into::<js_sys::Function>()
.map_err(|_| js_error("read not a function"))?;
let promise = read_fn
.call0(&self.reader)
.map_err(|_| js_error("read call failed"))?
.unchecked_into::<js_sys::Promise>();
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); return Ok(JsValue::NULL);
} }
let value = js_sys::Reflect::get(&result, &JsValue::from_str("value")) match self.stream.read_chunk().await? {
.map_err(|_| js_error("missing value"))?; Some(value) => Ok(js_sys::Uint8Array::from(&value[..]).into()),
Ok(js_sys::Uint8Array::new(&value).into()) None => {
self.stream.release();
self.finished = true;
Ok(JsValue::NULL)
}
}
} }
pub fn pipe_id(&self) -> u32 { pub fn pipe_id(&self) -> u32 {
@ -136,3 +118,9 @@ impl PipeReader {
self.description.clone() self.description.clone()
} }
} }
impl Drop for PipeReader {
fn drop(&mut self) {
self.stream.release();
}
}

View file

@ -9,6 +9,7 @@ use wasm_bindgen_futures::JsFuture;
use crate::error::js_error; use crate::error::js_error;
use crate::frame::parse_frame_value_with_limits; use crate::frame::parse_frame_value_with_limits;
use mtp_codec::{DecodeLimits, EncodeLimits, TypeMap}; use mtp_codec::{DecodeLimits, EncodeLimits, TypeMap};
use mtp_common::{FirstFrameDisposition, classify_first_frame};
const CLOSE_FRAME_LEN: u32 = u32::MAX; const CLOSE_FRAME_LEN: u32 = u32::MAX;
@ -25,12 +26,6 @@ pub(crate) fn log_stream_error_code(error: &JsValue, context: &str) {
let stream_error_code = js_sys::Reflect::get(error, &JsValue::from_str("streamErrorCode")) let stream_error_code = js_sys::Reflect::get(error, &JsValue::from_str("streamErrorCode"))
.ok() .ok()
.and_then(|v| v.as_f64()); .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 let message = error
.as_string() .as_string()
.or_else(|| { .or_else(|| {
@ -76,6 +71,233 @@ fn resolve_stream_readable(recv_stream: &JsValue) -> Result<JsValue, JsValue> {
} }
} }
#[derive(Clone)]
struct BrowserConnection {
inner: JsValue,
incoming_reader: Rc<RefCell<Option<JsValue>>>,
}
pub(crate) struct BrowserSendStream {
writer: JsValue,
}
pub(crate) struct BrowserRecvStream {
reader: JsValue,
}
impl BrowserConnection {
async fn connect(url: &str, cert_hashes: Option<Vec<String>>) -> Result<Self, JsValue> {
let constructor =
js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("WebTransport"))?
.dyn_into::<js_sys::Function>()
.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::<js_sys::Promise>()
.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<BrowserSendStream, JsValue> {
let create_stream = js_sys::Reflect::get(
&self.inner,
&JsValue::from_str("createUnidirectionalStream"),
)?
.dyn_into::<js_sys::Function>()
.map_err(|_| js_error("createUnidirectionalStream not a function"))?;
let stream_promise = create_stream
.call0(&self.inner)?
.dyn_into::<js_sys::Promise>()
.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::<js_sys::Function>()
.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<Option<BrowserRecvStream>, 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::<js_sys::Function>()
.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::<js_sys::Function>()
.map_err(|_| js_error("read not a function"))?;
let promise = read
.call0(&streams_reader)
.map_err(|_| js_error("read call failed"))?
.unchecked_into::<js_sys::Promise>();
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::<js_sys::Function>()
.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::<js_sys::Function>())
{
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::<js_sys::Function>()
.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::<js_sys::Promise>();
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::<js_sys::Function>()
.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::<js_sys::Promise>();
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::<js_sys::Function>()
.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<Option<Vec<u8>>, JsValue> {
let read = js_sys::Reflect::get(&self.reader, &JsValue::from_str("read"))
.map_err(|_| js_error("missing read"))?
.dyn_into::<js_sys::Function>()
.map_err(|_| js_error("read not a function"))?;
let promise = read
.call0(&self.reader)
.map_err(|_| js_error("read call failed"))?
.unchecked_into::<js_sys::Promise>();
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::<js_sys::Function>()
.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). /// 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) { pub(crate) fn release_writer_lock(writer: &JsValue) {
if let Ok(release) = js_sys::Reflect::get(writer, &JsValue::from_str("releaseLock")) if let Ok(release) = js_sys::Reflect::get(writer, &JsValue::from_str("releaseLock"))
@ -118,12 +340,10 @@ enum FrameOutcome {
*/ */
#[derive(Clone)] #[derive(Clone)]
pub struct WasmTransport { pub struct WasmTransport {
inner: JsValue, connection: BrowserConnection,
max_message_size: u32, max_message_size: u32,
/// Reader over `incoming_unidirectional_streams()` (a singleton stream of streams). /// Current incoming unidirectional stream, shared across handshake and receive loops.
streams_reader: Rc<RefCell<Option<JsValue>>>, stream_reader: Rc<RefCell<Option<BrowserRecvStream>>>,
/// Reader over the host's current uni-directional stream, if one is open.
stream_reader: Rc<RefCell<Option<JsValue>>>,
/// Bytes already read from the current stream but not yet consumed as a frame. /// Bytes already read from the current stream but not yet consumed as a frame.
buffer: Rc<RefCell<Vec<u8>>>, buffer: Rc<RefCell<Vec<u8>>>,
/// Set to `true` when `open_next_stream` succeeds; cleared after the first frame is parsed. /// Set to `true` when `open_next_stream` succeeds; cleared after the first frame is parsed.
@ -149,61 +369,14 @@ impl WasmTransport {
max_message_size: u32, max_message_size: u32,
configured_limits: Option<DecodeLimits>, configured_limits: Option<DecodeLimits>,
) -> Result<Self, JsValue> { ) -> Result<Self, JsValue> {
let ctor = js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("WebTransport"))? let connection = BrowserConnection::connect(url, cert_hashes).await?;
.dyn_into::<js_sys::Function>()
.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::<js_sys::Promise>()
.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 transport_limits = DecodeLimits::for_transport_message_size(max_message_size as u64);
let decode_limits = configured_limits let decode_limits = configured_limits
.map(|limits| restrict_decode_limits(limits, transport_limits)) .map(|limits| restrict_decode_limits(limits, transport_limits))
.unwrap_or(transport_limits); .unwrap_or(transport_limits);
Ok(Self { Ok(Self {
inner: transport, connection,
max_message_size, max_message_size,
streams_reader: Rc::new(RefCell::new(None)),
stream_reader: Rc::new(RefCell::new(None)), stream_reader: Rc::new(RefCell::new(None)),
buffer: Rc::new(RefCell::new(Vec::new())), buffer: Rc::new(RefCell::new(Vec::new())),
new_stream_frame: Rc::new(Cell::new(false)), new_stream_frame: Rc::new(Cell::new(false)),
@ -214,7 +387,7 @@ impl WasmTransport {
} }
pub fn inner(&self) -> &JsValue { pub fn inner(&self) -> &JsValue {
&self.inner &self.connection.inner
} }
pub fn set_type_map(&self, type_map: &TypeMap) { pub fn set_type_map(&self, type_map: &TypeMap) {
@ -250,154 +423,55 @@ impl WasmTransport {
// in accept_uni() until the authentication deadline. The bytes are // in accept_uni() until the authentication deadline. The bytes are
// already the canonical MTP self-framed value, so no extra stream // already the canonical MTP self-framed value, so no extra stream
// length prefix is added here. // length prefix is added here.
let create_stream = js_sys::Reflect::get( let mut stream = self.connection.open_uni().await?;
&self.inner, if let Err(e) = stream.write_all(frame).await {
&JsValue::from_str("createUnidirectionalStream"),
)?
.dyn_into::<js_sys::Function>()
.map_err(|_| js_error("createUnidirectionalStream not a function"))?;
let stream_promise = create_stream
.call0(&self.inner)?
.dyn_into::<js_sys::Promise>()
.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::<js_sys::Function>()
.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::<js_sys::Function>()
.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::<js_sys::Promise>()).await {
log_stream_error_code(&e, "send_frame write"); log_stream_error_code(&e, "send_frame write");
release_writer_lock(&writer_val); stream.release();
return Err(e); return Err(e);
} }
let close_fn = js_sys::Reflect::get(&writer_val, &JsValue::from_str("close")) if let Err(e) = stream.finish().await {
.map_err(|_| js_error("missing close"))?
.dyn_into::<js_sys::Function>()
.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::<js_sys::Promise>()).await {
// The frame was already written; do not retry it merely because // The frame was already written; do not retry it merely because
// FIN failed, as that would duplicate the MTP frame. // FIN failed, as that would duplicate the MTP frame.
log_stream_error_code(&e, "send_frame close"); log_stream_error_code(&e, "send_frame close");
} }
release_writer_lock(&writer_val); stream.release();
Ok(()) Ok(())
} }
/// Get (creating once) the reader over `incoming_unidirectional_streams()`.
fn ensure_streams_reader(&self) -> Result<JsValue, JsValue> {
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::<js_sys::Function>()
.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. /// Accept the next incoming uni-directional stream and make it current.
/// Returns `false` if the incoming-streams readable has ended. /// Returns `false` if the incoming-streams readable has ended.
async fn open_next_stream(&self) -> Result<bool, JsValue> { async fn open_next_stream(&self) -> Result<bool, JsValue> {
let streams_reader = self.ensure_streams_reader()?; let Some(stream) = self.connection.accept_uni().await? else {
return Ok(false);
let read_fn = js_sys::Reflect::get(&streams_reader, &JsValue::from_str("read"))
.map_err(|_| js_error("missing read"))?
.dyn_into::<js_sys::Function>()
.map_err(|_| js_error("read not a function"))?;
let promise = read_fn
.call0(&streams_reader)
.map_err(|_| js_error("read call failed"))?
.unchecked_into::<js_sys::Promise>();
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)));
}
}; };
let done = js_sys::Reflect::get(&result, &JsValue::from_str("done")) *self.stream_reader.borrow_mut() = Some(stream);
.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::<js_sys::Function>()
.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); self.new_stream_frame.set(true);
Ok(true) Ok(true)
} }
/// Read one chunk from the current stream. `Ok(None)` means the stream ended. /// Read one chunk from the current stream. `Ok(None)` means the stream ended.
async fn read_chunk(&self) -> Result<Option<Vec<u8>>, JsValue> { async fn read_chunk(&self) -> Result<Option<Vec<u8>>, JsValue> {
let reader = match self.stream_reader.borrow().clone() { let mut stream = match self.stream_reader.borrow_mut().take() {
Some(r) => r, Some(stream) => stream,
None => return Ok(None), None => return Ok(None),
}; };
let result = match stream.read_chunk().await {
let read_fn = js_sys::Reflect::get(&reader, &JsValue::from_str("read")) Ok(result) => result,
.map_err(|_| js_error("missing read"))?
.dyn_into::<js_sys::Function>()
.map_err(|_| js_error("read not a function"))?;
let promise = read_fn
.call0(&reader)
.map_err(|_| js_error("read call failed"))?
.unchecked_into::<js_sys::Promise>();
let result = match JsFuture::from(promise).await {
Ok(r) => r,
Err(e) => { Err(e) => {
log_stream_error_code(&e, "read_chunk"); log_stream_error_code(&e, "read_chunk");
stream.release();
return Err(js_error(format!("read failed: {:?}", e))); return Err(js_error(format!("read failed: {:?}", e)));
} }
}; };
if result.is_some() {
let done = js_sys::Reflect::get(&result, &JsValue::from_str("done")) *self.stream_reader.borrow_mut() = Some(stream);
.ok() } else {
.and_then(|v| v.as_bool()) stream.release();
.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. /// Try to pull one complete frame out of the buffer without reading more.
@ -459,9 +533,7 @@ impl WasmTransport {
} }
None => { None => {
// Stream finished; release the reader's lock to avoid a spurious cancel. // Stream finished; release the reader's lock to avoid a spurious cancel.
if let Some(reader) = self.stream_reader.borrow_mut().take() { // `read_chunk` releases the raw stream lock on clean FIN.
release_reader_lock(&reader);
}
// A frame is never allowed to span stream boundaries. The // A frame is never allowed to span stream boundaries. The
// native persistent-stream sender packs frames on one // native persistent-stream sender packs frames on one
// stream, while the WASM sender uses one stream per frame; // stream, while the WASM sender uses one stream per frame;
@ -520,15 +592,17 @@ impl WasmTransport {
/// Pipe-aware receive loop. Identical to `receive_loop` but detects /// Pipe-aware receive loop. Identical to `receive_loop` but detects
/// `PipeRequest` as the first frame on a new incoming stream and routes /// `PipeRequest` as the first frame on a new incoming stream and routes
/// the stream to `on_pipe` instead of `on_message`. /// the stream to `on_pipe` instead of `on_message`.
pub async fn receive_loop_with_pipes<F, G, H>( pub async fn receive_loop_with_pipes<F, G, H, I>(
&self, &self,
mut on_message: F, mut on_message: F,
mut on_error: H, mut on_error: H,
mut on_pipe: G, mut on_pipe: G,
mut pipe_is_expected: I,
) where ) where
F: FnMut(JsValue), F: FnMut(JsValue),
G: FnMut(crate::pipe::PipeReader), G: FnMut(crate::pipe::PipeReader),
H: FnMut(JsValue), H: FnMut(JsValue),
I: FnMut(u32) -> bool,
{ {
loop { loop {
match self.next_frame(self.max_message_size).await { match self.next_frame(self.max_message_size).await {
@ -550,36 +624,44 @@ impl WasmTransport {
if is_first { if is_first {
self.new_stream_frame.set(false); self.new_stream_frame.set(false);
if let Some(comm) = comm.as_ref() if let Some(comm) = comm.as_ref() {
&& Some(comm.get_type()) == pipe_request_type let is_pipe_request = Some(comm.get_type()) == pipe_request_type;
{ let pipe_id = comm.id().filter(|id| *id != 0);
let Some(pipe_id) = comm.id().filter(|id| *id != 0) else { let is_expected =
on_error(JsValue::from_str( is_pipe_request && pipe_id.is_some_and(&mut pipe_is_expected);
"PipeRequest frame must contain a non-zero id", let disposition =
)); match classify_first_frame(is_pipe_request, comm.id(), is_expected)
self.close(); {
break; Ok(disposition) => disposition,
}; Err(error) => {
let description = comm on_error(JsValue::from_str(&error.to_string()));
.get_str(mtp_codec::DataType::Description) self.close();
.unwrap_or("") break;
.to_string(); }
};
let pending = { if let FirstFrameDisposition::Pipe(pipe_id) = disposition {
let mut buf = self.buffer.borrow_mut(); let description = comm
std::mem::take(&mut *buf) .get_str(mtp_codec::DataType::Description)
}; .unwrap_or("")
.to_string();
if let Some(reader) = self.stream_reader.borrow_mut().take() { let pending = {
let pipe_reader = crate::pipe::PipeReader::new( let mut buf = self.buffer.borrow_mut();
reader, std::mem::take(&mut *buf)
pipe_id, };
description,
pending, if let Some(reader) = self.stream_reader.borrow_mut().take() {
); let pipe_reader = crate::pipe::PipeReader::new(
on_pipe(pipe_reader); reader,
pipe_id,
description,
pending,
);
on_pipe(pipe_reader);
}
continue;
} }
continue;
} }
} }
@ -635,25 +717,7 @@ impl WasmTransport {
description: &str, description: &str,
) -> Result<crate::pipe::PipeWriter, JsValue> { ) -> Result<crate::pipe::PipeWriter, JsValue> {
let _send_guard = self.send_lock.lock().await; let _send_guard = self.send_lock.lock().await;
let create_stream = js_sys::Reflect::get( let mut stream = self.connection.open_uni().await?;
&self.inner,
&JsValue::from_str("createUnidirectionalStream"),
)?
.dyn_into::<js_sys::Function>()
.map_err(|_| js_error("createUnidirectionalStream not a function"))?;
let stream_promise = create_stream
.call0(&self.inner)?
.dyn_into::<js_sys::Promise>()
.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::<js_sys::Function>()
.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 type_map = self.type_map();
let request = mtp_codec::CommunicationValue::new_with_type_map( let request = mtp_codec::CommunicationValue::new_with_type_map(
@ -669,37 +733,21 @@ impl WasmTransport {
.to_bytes() .to_bytes()
.map_err(|e| js_error(format!("encode failed: {}", e)))?; .map_err(|e| js_error(format!("encode failed: {}", e)))?;
let chunk = js_sys::Uint8Array::from(&frame_bytes[..]); if let Err(e) = stream.write_all(&frame_bytes).await {
let write_fn = js_sys::Reflect::get(&writer_val, &JsValue::from_str("write"))
.map_err(|_| js_error("missing write"))?
.dyn_into::<js_sys::Function>()
.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::<js_sys::Promise>()).await {
log_stream_error_code(&e, "open_pipe write"); log_stream_error_code(&e, "open_pipe write");
release_writer_lock(&writer_val); stream.release();
return Err(e); return Err(e);
} }
Ok(crate::pipe::PipeWriter::new(writer_val, pipe_id)) Ok(crate::pipe::PipeWriter::new(stream, pipe_id))
} }
pub fn close(&self) { pub fn close(&self) {
// Release reader locks before closing so they aren't treated as cancels. // Release reader locks before closing so they aren't treated as cancels.
if let Some(reader) = self.stream_reader.borrow_mut().take() { if let Some(reader) = self.stream_reader.borrow_mut().take() {
release_reader_lock(&reader); reader.release();
}
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::<js_sys::Function>())
{
let _ = close.call1(&self.inner, &js_sys::Object::new());
} }
self.connection.close();
} }
} }