smb_server_transport/
tcp.rs1use async_trait::async_trait;
11use std::io;
12use std::rc::Rc;
13use tokio_uring::net::TcpStream;
14
15use crate::{Frame, FrameSink, FrameSource, Transport, TransportError};
16
17mod nb_type {
19 pub const SESSION_MESSAGE: u8 = 0x00;
21 pub const SESSION_REQUEST: u8 = 0x81;
23 pub const POSITIVE_RESPONSE: u8 = 0x82;
25}
26
27const MAX_FRAME: usize = 0x20_0000;
29const NBSS_HEADER_LEN: usize = 4;
31const READ_CHUNK: usize = 64 * 1024;
33
34struct NbssReader {
36 stream: Rc<TcpStream>,
37 carry: Vec<u8>,
39}
40
41impl std::fmt::Debug for NbssReader {
42 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
43 f.debug_struct("NbssReader").field("carry", &self.carry.len()).finish()
44 }
45}
46
47impl NbssReader {
48 fn new(stream: Rc<TcpStream>) -> Self {
49 Self { stream, carry: Vec::new() }
50 }
51
52 async fn fill(&mut self) -> Result<bool, TransportError> {
54 let (res, chunk) = self.stream.read(vec![0u8; READ_CHUNK]).await;
55 let n = res?;
56 if n == 0 {
57 return Ok(false);
58 }
59 self.carry.extend_from_slice(&chunk[..n]);
60 Ok(true)
61 }
62
63 async fn take(&mut self, n: usize) -> Result<Option<Vec<u8>>, TransportError> {
65 while self.carry.len() < n {
66 if !self.fill().await? {
67 return Ok(None);
68 }
69 }
70 let rest = self.carry.split_off(n);
71 Ok(Some(std::mem::replace(&mut self.carry, rest)))
72 }
73
74 async fn next_frame(
78 &mut self,
79 mut answer: Option<&Rc<TcpStream>>,
80 ) -> Result<Option<Frame>, TransportError> {
81 loop {
82 let hdr = match self.take(NBSS_HEADER_LEN).await? {
83 Some(h) => h,
84 None => return Ok(None),
85 };
86 let len = decode_nbss_len(&hdr);
87 match hdr[0] {
88 nb_type::SESSION_MESSAGE => {
89 if len > MAX_FRAME {
90 return Err(io::Error::new(
91 io::ErrorKind::InvalidData,
92 "oversized NBSS frame",
93 )
94 .into());
95 }
96 return Ok(self.take(len).await?.map(Frame));
97 }
98 nb_type::SESSION_REQUEST => {
99 if len > 0 && self.take(len).await?.is_none() {
100 return Ok(None);
101 }
102 if let Some(s) = answer.take() {
103 write_nbss(s, &[nb_type::POSITIVE_RESPONSE, 0, 0, 0]).await?;
104 }
105 }
106 _ => {
107 if len > 0 && self.take(len).await?.is_none() {
109 return Ok(None);
110 }
111 }
112 }
113 }
114 }
115}
116
117#[cfg_attr(dylint_lib = "no_magic_numbers", allow(no_magic_numbers))]
124fn decode_nbss_len(hdr: &[u8]) -> usize {
125 ((hdr[1] as usize) << 16) | ((hdr[2] as usize) << 8) | hdr[3] as usize
126}
127
128#[cfg_attr(dylint_lib = "no_magic_numbers", allow(no_magic_numbers))]
130fn encode_nbss_len(len: usize) -> [u8; 3] {
131 [(len >> 16) as u8, (len >> 8) as u8, (len & 0xff) as u8]
132}
133
134async fn write_nbss(stream: &TcpStream, data: &[u8]) -> Result<(), TransportError> {
136 let len = encode_nbss_len(data.len());
139 let mut frame = Vec::with_capacity(NBSS_HEADER_LEN + data.len());
140 frame.push(nb_type::SESSION_MESSAGE);
141 frame.extend_from_slice(&len);
142 frame.extend_from_slice(data);
143 let (res, _buf) = stream.write_all(frame).await;
144 res.map_err(Into::into)
145}
146
147pub struct TcpTransport {
149 reader: NbssReader,
150 stream: Rc<TcpStream>,
151 peer: String,
152}
153
154impl std::fmt::Debug for TcpTransport {
155 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
156 f.debug_struct("TcpTransport").field("peer", &self.peer).finish()
157 }
158}
159
160impl TcpTransport {
161 pub fn new(stream: TcpStream, peer: String) -> Self {
164 let stream = Rc::new(stream);
165 Self {
166 reader: NbssReader::new(stream.clone()),
167 stream,
168 peer,
169 }
170 }
171}
172
173#[async_trait(?Send)]
174impl Transport for TcpTransport {
175 async fn recv(&mut self) -> Result<Option<Frame>, TransportError> {
176 self.reader.next_frame(Some(&self.stream)).await
177 }
178
179 async fn send(&mut self, data: &[u8]) -> Result<(), TransportError> {
180 write_nbss(&self.stream, data).await
181 }
182
183 fn peer(&self) -> String {
184 self.peer.clone()
185 }
186
187 fn split(self: Box<Self>) -> (Box<dyn FrameSource>, Box<dyn FrameSink>) {
188 (
189 Box::new(TcpSource {
190 reader: self.reader,
191 peer: self.peer,
192 }),
193 Box::new(TcpSink {
194 stream: self.stream,
195 }),
196 )
197 }
198}
199
200pub struct TcpSource {
202 reader: NbssReader,
203 peer: String,
204}
205
206impl std::fmt::Debug for TcpSource {
207 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
208 f.debug_struct("TcpSource").field("peer", &self.peer).finish()
209 }
210}
211
212#[async_trait(?Send)]
213impl FrameSource for TcpSource {
214 async fn recv(&mut self) -> Result<Option<Frame>, TransportError> {
215 self.reader.next_frame(None).await
217 }
218
219 fn peer(&self) -> String {
220 self.peer.clone()
221 }
222}
223
224pub struct TcpSink {
226 stream: Rc<TcpStream>,
227}
228
229impl std::fmt::Debug for TcpSink {
230 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
231 f.debug_struct("TcpSink").finish()
232 }
233}
234
235#[async_trait(?Send)]
236impl FrameSink for TcpSink {
237 async fn send(&mut self, data: &[u8]) -> Result<(), TransportError> {
238 write_nbss(&self.stream, data).await
239 }
240}
241
242#[cfg(test)]
243mod nbss_len_tests {
244 use super::{decode_nbss_len, encode_nbss_len, MAX_FRAME};
245
246 #[test]
250 fn round_trips_across_full_range() {
251 for len in [0usize, 1, 63, 64 * 1024, 128 * 1024, 256 * 1024, MAX_FRAME] {
252 let enc = encode_nbss_len(len);
253 let hdr = [0u8, enc[0], enc[1], enc[2]];
254 assert_eq!(decode_nbss_len(&hdr), len, "round-trip failed for {len}");
255 }
256 }
257
258 #[test]
260 fn decodes_high_byte() {
261 let enc = encode_nbss_len(256 * 1024);
262 assert_eq!(enc, [0x04, 0x00, 0x00]);
263 assert_eq!(decode_nbss_len(&[0x00, 0x04, 0x00, 0x00]), 256 * 1024);
264 }
265}