Skip to main content

smb_server_transport/
mem.rs

1//! In-memory transport for testing protocol layers without sockets.
2//!
3//! A [`mem::Duplex`] pair behaves like a crossed wire: frames written into
4//! one side are returned by the other side's [`Transport::recv`].
5
6use async_trait::async_trait;
7use tokio::sync::mpsc;
8
9use crate::{Frame, FrameSink, FrameSource, Transport, TransportError};
10
11/// One end of an in-memory transport pair.
12#[derive(Debug)]
13pub struct MemTransport {
14    rx: mpsc::Receiver<Vec<u8>>,
15    tx: mpsc::Sender<Vec<u8>>,
16    peer: String,
17}
18
19/// Create two connected transports. Frames sent on one are received by the
20/// other; dropping a sender closes that direction.
21pub fn duplex(buffer: usize) -> (MemTransport, MemTransport) {
22    let (tx_a, rx_a) = mpsc::channel(buffer);
23    let (tx_b, rx_b) = mpsc::channel(buffer);
24    (
25        MemTransport { rx: rx_a, tx: tx_b, peer: "mem:a".into() },
26        MemTransport { rx: rx_b, tx: tx_a, peer: "mem:b".into() },
27    )
28}
29
30#[async_trait(?Send)]
31impl Transport for MemTransport {
32    async fn recv(&mut self) -> Result<Option<Frame>, TransportError> {
33        Ok(self.rx.recv().await.map(Frame))
34    }
35
36    async fn send(&mut self, data: &[u8]) -> Result<(), TransportError> {
37        self.tx
38            .send(data.to_vec())
39            .await
40            .map_err(|_| TransportError::Io(std::io::Error::new(
41                std::io::ErrorKind::BrokenPipe,
42                "peer dropped",
43            )))
44    }
45
46    fn peer(&self) -> String {
47        self.peer.clone()
48    }
49
50    fn split(self: Box<Self>) -> (Box<dyn FrameSource>, Box<dyn FrameSink>) {
51        (
52            Box::new(MemSource { rx: self.rx, peer: self.peer }),
53            Box::new(MemSink { tx: self.tx }),
54        )
55    }
56}
57
58/// Read half of a split [`MemTransport`].
59#[derive(Debug)]
60pub struct MemSource {
61    rx: mpsc::Receiver<Vec<u8>>,
62    peer: String,
63}
64
65#[async_trait(?Send)]
66impl FrameSource for MemSource {
67    async fn recv(&mut self) -> Result<Option<Frame>, TransportError> {
68        Ok(self.rx.recv().await.map(Frame))
69    }
70
71    fn peer(&self) -> String {
72        self.peer.clone()
73    }
74}
75
76/// Write half of a split [`MemTransport`].
77#[derive(Debug)]
78pub struct MemSink {
79    tx: mpsc::Sender<Vec<u8>>,
80}
81
82#[async_trait(?Send)]
83impl FrameSink for MemSink {
84    async fn send(&mut self, data: &[u8]) -> Result<(), TransportError> {
85        self.tx
86            .send(data.to_vec())
87            .await
88            .map_err(|_| TransportError::Io(std::io::Error::new(
89                std::io::ErrorKind::BrokenPipe,
90                "peer dropped",
91            )))
92    }
93}