#![cfg(feature = "pipes")] use async_trait::async_trait; use mtp_codec::CommunicationValue; use mtp_common::CommunicationError; use mtp_transport::{ GenericReceiver, GenericSender, Policy, SendMode, TransportConnection, TransportEvent, TransportRecvStream, TransportSendStream, }; use std::sync::Arc; use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, DuplexStream, duplex}; use tokio::sync::{Mutex, mpsc}; struct MockSendStream { inner: DuplexStream, } impl AsyncWrite for MockSendStream { fn poll_write( mut self: std::pin::Pin<&mut Self>, cx: &mut std::task::Context<'_>, buf: &[u8], ) -> std::task::Poll> { std::pin::Pin::new(&mut self.inner).poll_write(cx, buf) } fn poll_flush( mut self: std::pin::Pin<&mut Self>, cx: &mut std::task::Context<'_>, ) -> std::task::Poll> { std::pin::Pin::new(&mut self.inner).poll_flush(cx) } fn poll_shutdown( mut self: std::pin::Pin<&mut Self>, cx: &mut std::task::Context<'_>, ) -> std::task::Poll> { std::pin::Pin::new(&mut self.inner).poll_shutdown(cx) } } #[async_trait] impl TransportSendStream for MockSendStream { async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError> { AsyncWriteExt::write_all(&mut self.inner, buf) .await .map_err(|_| CommunicationError::StreamError) } async fn finish(&mut self) -> Result<(), CommunicationError> { self.inner .shutdown() .await .map_err(|_| CommunicationError::StreamError) } fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> { Ok(()) } } struct MockRecvStream { inner: DuplexStream, } impl AsyncRead for MockRecvStream { fn poll_read( mut self: std::pin::Pin<&mut Self>, cx: &mut std::task::Context<'_>, buf: &mut tokio::io::ReadBuf<'_>, ) -> std::task::Poll> { std::pin::Pin::new(&mut self.inner).poll_read(cx, buf) } } #[async_trait] impl TransportRecvStream for MockRecvStream { async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError> { AsyncReadExt::read_exact(&mut self.inner, buf) .await .map(|_| ()) .map_err(|_| CommunicationError::StreamError) } async fn read_chunk(&mut self, max: usize) -> Result>, CommunicationError> { let mut buf = vec![0u8; max]; match AsyncReadExt::read(&mut self.inner, &mut buf).await { Ok(0) => Ok(None), Ok(n) => { buf.truncate(n); Ok(Some(buf)) } Err(_) => Err(CommunicationError::StreamError), } } fn stop(self, _code: u32) -> Result<(), CommunicationError> { Ok(()) } } #[derive(Clone)] struct MockTransportConnection { pair_tx: mpsc::Sender, pair_rx: Arc>>, } impl MockTransportConnection { fn pair() -> (Self, Self) { let (tx_a, rx_a) = mpsc::channel(16); let (tx_b, rx_b) = mpsc::channel(16); ( Self { pair_tx: tx_a, pair_rx: Arc::new(Mutex::new(rx_b)), }, Self { pair_tx: tx_b, pair_rx: Arc::new(Mutex::new(rx_a)), }, ) } } #[async_trait] impl TransportConnection for MockTransportConnection { type SendStream = MockSendStream; type RecvStream = MockRecvStream; async fn open_uni(&self) -> Result { let (local, remote) = duplex(65536); self.pair_tx .send(remote) .await .map_err(|_| CommunicationError::StreamError)?; Ok(MockSendStream { inner: local }) } async fn accept_uni(&self) -> Result { let remote = self .pair_rx .lock() .await .recv() .await .ok_or(CommunicationError::StreamClosed)?; Ok(MockRecvStream { inner: remote }) } fn close_reason(&self) -> Option { None } fn close(&self, _code: u32, _reason: &[u8]) {} } async fn mock_connected_pair() -> (MockTransportConnection, MockTransportConnection) { MockTransportConnection::pair() } #[tokio::test] async fn test_open_pipe_and_receive_reader() -> Result<(), Box> { let (conn_a, conn_b) = mock_connected_pair().await; let policy = Arc::new(Policy::default()); let sender = GenericSender::new(conn_a, policy.clone()); let receiver = GenericReceiver::new(conn_b, policy); receiver.expect_pipe(42)?; let pipe_writer = sender.open_pipe(42, "test-pipe").await?; let pipe_reader = receiver.receive_pipe().await?; assert_eq!(pipe_reader.pipe_id(), 42); assert_eq!(pipe_reader.description(), "test-pipe"); drop(pipe_writer); Ok(()) } #[tokio::test] async fn test_pipe_raw_data_roundtrip() -> Result<(), Box> { let (conn_a, conn_b) = mock_connected_pair().await; let policy = Arc::new(Policy::default()); 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"; AsyncWriteExt::write_all(&mut pipe_writer, data).await?; pipe_writer.finish_async().await?; let mut pipe_reader = receiver.receive_pipe().await?; let mut buf = vec![0u8; data.len()]; AsyncReadExt::read_exact(&mut pipe_reader, &mut buf).await?; assert_eq!(&buf, data); Ok(()) } #[tokio::test] async fn test_pipe_large_payload() -> Result<(), Box> { let (conn_a, conn_b) = mock_connected_pair().await; let policy = Arc::new(Policy::default()); 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(); let data_clone = data.clone(); let write_handle = tokio::spawn(async move { AsyncWriteExt::write_all(&mut pipe_writer, &data_clone) .await .map_err(|_| CommunicationError::StreamError)?; pipe_writer.finish_async().await }); let mut pipe_reader = receiver.receive_pipe().await?; let mut buf = Vec::new(); AsyncReadExt::read_to_end(&mut pipe_reader, &mut buf).await?; assert_eq!(buf, data); write_handle.await??; Ok(()) } #[tokio::test] async fn test_receive_event_dispatches_pipe() -> Result<(), Box> { let (conn_a, conn_b) = mock_connected_pair().await; let policy = Arc::new(Policy::default()); let sender = GenericSender::new(conn_a, policy.clone()); let receiver = GenericReceiver::new(conn_b, policy); receiver.expect_pipe(99)?; let mut pipe_writer = sender.open_pipe(99, "event-pipe").await?; match receiver.receive_event().await? { TransportEvent::Pipe(mut reader) => { assert_eq!(reader.pipe_id(), 99); assert_eq!(reader.description(), "event-pipe"); let data = b"event dispatch test"; AsyncWriteExt::write_all(&mut pipe_writer, data).await?; pipe_writer.finish_async().await?; let mut buf = vec![0u8; data.len()]; AsyncReadExt::read_exact(&mut reader, &mut buf).await?; assert_eq!(&buf, data); } TransportEvent::Message(_) => panic!("expected Pipe event, got Message"), } Ok(()) } #[tokio::test] async fn test_try_receive_pipe_returns_none_when_empty() -> Result<(), Box> { let (conn_a, conn_b) = mock_connected_pair().await; let policy = Arc::new(Policy::default()); let _sender = GenericSender::new(conn_a, policy.clone()); let receiver = GenericReceiver::new(conn_b, policy); let result = receiver.try_receive_pipe()?; assert!(result.is_none()); Ok(()) } #[tokio::test] async fn test_regular_messages_still_work_alongside_pipes() -> Result<(), Box> { let (conn_a, conn_b) = mock_connected_pair().await; let policy = Arc::new(Policy::default().with_send_mode(SendMode::SingleStreamPerMessage)); 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 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?; let pipe_reader = receiver.receive_pipe().await?; assert_eq!(pipe_reader.pipe_id(), 1); Ok(()) } #[tokio::test] async fn test_multiple_pipes() -> Result<(), Box> { let (conn_a, conn_b) = mock_connected_pair().await; let policy = Arc::new(Policy::default()); 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?; let r1 = receiver.receive_pipe().await?; assert_eq!(r1.pipe_id(), 10); let r2 = receiver.receive_pipe().await?; assert_eq!(r2.pipe_id(), 20); let data1 = b"pipe-one-data"; AsyncWriteExt::write_all(&mut pw1, data1).await?; pw1.finish_async().await?; let data2 = b"pipe-two-data"; AsyncWriteExt::write_all(&mut pw2, data2).await?; pw2.finish_async().await?; let mut buf1 = vec![0u8; data1.len()]; let mut reader1 = r1; AsyncReadExt::read_exact(&mut reader1, &mut buf1).await?; assert_eq!(&buf1, data1); let mut buf2 = vec![0u8; data2.len()]; let mut reader2 = r2; AsyncReadExt::read_exact(&mut reader2, &mut buf2).await?; assert_eq!(&buf2, data2); Ok(()) }