1use std::{error::Error, fmt, io};
10
11use subc_protocol::{
12 decode_header, DecodeError, Frame, FROZEN_PREFIX_LEN, HEADER_LEN, MAX_FRAME_BODY_LEN,
13 PROTOCOL_VERSION,
14};
15use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
16
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
19pub enum ReadStage {
20 Header,
21 Body,
22}
23
24#[derive(Debug)]
26pub enum FrameIoError {
27 Io(io::Error),
28 DecodeHeader(DecodeError),
29 BodyTooLarge {
30 len: u32,
31 max: u32,
32 },
33 UnexpectedEof {
34 stage: ReadStage,
35 expected: usize,
36 actual: usize,
37 },
38 BodyLengthMismatch {
39 header_len: u32,
40 body_len: usize,
41 },
42}
43
44pub async fn read_frame<R>(reader: &mut R) -> Result<Option<Frame>, FrameIoError>
50where
51 R: AsyncRead + Unpin,
52{
53 let mut prefix = [0u8; FROZEN_PREFIX_LEN];
54 if !read_exact_or_clean_eof(reader, &mut prefix, ReadStage::Header).await? {
55 return Ok(None);
56 }
57 let ver = prefix[4];
58 if ver != PROTOCOL_VERSION {
59 return Err(FrameIoError::DecodeHeader(
60 DecodeError::UnsupportedVersion { ver },
61 ));
62 }
63
64 let mut header_bytes = [0u8; HEADER_LEN];
65 header_bytes[..FROZEN_PREFIX_LEN].copy_from_slice(&prefix);
66 read_exact_or_unexpected_eof(
67 reader,
68 &mut header_bytes[FROZEN_PREFIX_LEN..],
69 ReadStage::Header,
70 )
71 .await?;
72
73 let header = decode_header(&header_bytes).map_err(FrameIoError::DecodeHeader)?;
74 if header.len > MAX_FRAME_BODY_LEN {
75 return Err(FrameIoError::BodyTooLarge {
76 len: header.len,
77 max: MAX_FRAME_BODY_LEN,
78 });
79 }
80 let body_len = header.len as usize;
81 let mut body = vec![0u8; body_len];
82 if body_len > 0 {
83 read_exact_or_unexpected_eof(reader, &mut body, ReadStage::Body).await?;
84 }
85
86 Ok(Some(Frame::from_wire(header, body)))
87}
88
89pub async fn write_frame<W>(writer: &mut W, frame: &Frame) -> Result<(), FrameIoError>
107where
108 W: AsyncWrite + Unpin,
109{
110 if frame.header.len as usize != frame.body.len() {
111 return Err(FrameIoError::BodyLengthMismatch {
112 header_len: frame.header.len,
113 body_len: frame.body.len(),
114 });
115 }
116
117 let header = frame.header.encode();
118 if frame.body.is_empty() {
119 return writer.write_all(&header).await.map_err(FrameIoError::Io);
120 }
121
122 let mut joined = Vec::with_capacity(header.len() + frame.body.len());
123 joined.extend_from_slice(&header);
124 joined.extend_from_slice(&frame.body);
125 writer.write_all(&joined).await.map_err(FrameIoError::Io)
126}
127
128async fn read_exact_or_clean_eof<R>(
129 reader: &mut R,
130 buf: &mut [u8],
131 stage: ReadStage,
132) -> Result<bool, FrameIoError>
133where
134 R: AsyncRead + Unpin,
135{
136 let mut actual = 0;
137 while actual < buf.len() {
138 let n = reader
139 .read(&mut buf[actual..])
140 .await
141 .map_err(FrameIoError::Io)?;
142 if n == 0 {
143 if actual == 0 {
144 return Ok(false);
145 }
146 return Err(FrameIoError::UnexpectedEof {
147 stage,
148 expected: buf.len(),
149 actual,
150 });
151 }
152 actual += n;
153 }
154 Ok(true)
155}
156
157async fn read_exact_or_unexpected_eof<R>(
158 reader: &mut R,
159 buf: &mut [u8],
160 stage: ReadStage,
161) -> Result<(), FrameIoError>
162where
163 R: AsyncRead + Unpin,
164{
165 let mut actual = 0;
166 while actual < buf.len() {
167 let n = reader
168 .read(&mut buf[actual..])
169 .await
170 .map_err(FrameIoError::Io)?;
171 if n == 0 {
172 return Err(FrameIoError::UnexpectedEof {
173 stage,
174 expected: buf.len(),
175 actual,
176 });
177 }
178 actual += n;
179 }
180 Ok(())
181}
182
183impl fmt::Display for FrameIoError {
184 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
185 match self {
186 Self::Io(err) => write!(f, "frame I/O error: {err}"),
187 Self::DecodeHeader(err) => write!(f, "invalid envelope header: {err}"),
188 Self::BodyTooLarge { len, max } => {
189 write!(f, "frame body length {len} exceeds max {max}")
190 }
191 Self::UnexpectedEof {
192 stage,
193 expected,
194 actual,
195 } => write!(
196 f,
197 "unexpected EOF while reading {stage:?}: expected {expected} bytes, got {actual}"
198 ),
199 Self::BodyLengthMismatch {
200 header_len,
201 body_len,
202 } => write!(
203 f,
204 "frame header len ({header_len}) does not match body length ({body_len})"
205 ),
206 }
207 }
208}
209
210impl Error for FrameIoError {
211 fn source(&self) -> Option<&(dyn Error + 'static)> {
212 match self {
213 Self::Io(err) => Some(err),
214 Self::DecodeHeader(err) => Some(err),
215 Self::UnexpectedEof { .. }
216 | Self::BodyTooLarge { .. }
217 | Self::BodyLengthMismatch { .. } => None,
218 }
219 }
220}
221
222impl From<io::Error> for FrameIoError {
223 fn from(err: io::Error) -> Self {
224 Self::Io(err)
225 }
226}
227
228#[cfg(test)]
229mod tests {
230 use super::*;
231 use subc_protocol::{Flags, FrameType, Priority, PROTOCOL_VERSION};
232 use tokio::io::{duplex, AsyncWriteExt};
233
234 fn test_frame(channel: u16, corr: u64, body: &[u8]) -> Frame {
235 Frame::build(
236 FrameType::Request,
237 Flags::new(true, Priority::Interactive, false),
238 channel,
239 1,
240 corr,
241 body.to_vec(),
242 )
243 .unwrap()
244 }
245
246 #[derive(Default)]
251 struct WriteCounter {
252 writes: Vec<usize>,
253 bytes: Vec<u8>,
254 }
255
256 impl AsyncWrite for WriteCounter {
257 fn poll_write(
258 mut self: std::pin::Pin<&mut Self>,
259 _cx: &mut std::task::Context<'_>,
260 buf: &[u8],
261 ) -> std::task::Poll<io::Result<usize>> {
262 self.writes.push(buf.len());
263 self.bytes.extend_from_slice(buf);
264 std::task::Poll::Ready(Ok(buf.len()))
265 }
266
267 fn poll_flush(
268 self: std::pin::Pin<&mut Self>,
269 _cx: &mut std::task::Context<'_>,
270 ) -> std::task::Poll<io::Result<()>> {
271 std::task::Poll::Ready(Ok(()))
272 }
273
274 fn poll_shutdown(
275 self: std::pin::Pin<&mut Self>,
276 _cx: &mut std::task::Context<'_>,
277 ) -> std::task::Poll<io::Result<()>> {
278 std::task::Poll::Ready(Ok(()))
279 }
280 }
281
282 #[tokio::test]
291 async fn a_frame_with_a_body_reaches_the_socket_as_one_write() {
292 let mut writer = WriteCounter::default();
293 let frame = test_frame(3, 11, &vec![0xABu8; 16 * 1024]);
294
295 write_frame(&mut writer, &frame).await.unwrap();
296
297 assert_eq!(
298 writer.writes.len(),
299 1,
300 "header and body must be one write, got segments {:?}",
301 writer.writes
302 );
303 assert_eq!(writer.writes[0], HEADER_LEN + frame.body.len());
304
305 let mut expected = frame.header.encode().to_vec();
308 expected.extend_from_slice(&frame.body);
309 assert_eq!(writer.bytes, expected);
310 }
311
312 #[tokio::test]
315 async fn a_bodyless_frame_writes_only_its_header() {
316 let mut writer = WriteCounter::default();
317 let frame = test_frame(4, 12, b"");
318
319 write_frame(&mut writer, &frame).await.unwrap();
320
321 assert_eq!(writer.writes, vec![HEADER_LEN]);
322 }
323
324 #[tokio::test]
325 async fn read_write_round_trip_preserves_opaque_body() {
326 let (mut client, mut server) = duplex(128);
327 let frame = test_frame(7, 42, b"opaque\0json? no parse");
328 let expected = frame.clone();
329
330 let writer = tokio::spawn(async move { write_frame(&mut client, &frame).await });
331 let read = read_frame(&mut server).await.unwrap().unwrap();
332
333 writer.await.unwrap().unwrap();
334 assert_eq!(read, expected);
335 }
336
337 #[tokio::test]
338 async fn partial_header_and_body_are_assembled() {
339 let (mut client, mut server) = duplex(128);
340 let frame = test_frame(2, 99, b"chunked-body");
341 let mut bytes = frame.header.encode().to_vec();
342 bytes.extend_from_slice(&frame.body);
343 let expected = frame.clone();
344
345 let writer = tokio::spawn(async move {
346 client.write_all(&bytes[..3]).await.unwrap();
347 client.write_all(&bytes[3..10]).await.unwrap();
348 client.write_all(&bytes[10..]).await.unwrap();
349 });
350
351 let read = read_frame(&mut server).await.unwrap().unwrap();
352 writer.await.unwrap();
353 assert_eq!(read, expected);
354 }
355
356 #[tokio::test]
357 async fn clean_eof_before_header_returns_none() {
358 let (client, mut server) = duplex(16);
359 drop(client);
360
361 assert!(read_frame(&mut server).await.unwrap().is_none());
362 }
363
364 #[tokio::test]
365 async fn stale_v1_pure_header_is_rejected_from_prefix_without_waiting() {
366 let (mut client, mut server) = duplex(64);
367 let mut stale_header = [0u8; 17];
368 stale_header[4] = 1;
369 stale_header[5] = FrameType::Ping as u8;
370 client.write_all(&stale_header).await.unwrap();
371
372 let err = tokio::time::timeout(
373 std::time::Duration::from_millis(100),
374 read_frame(&mut server),
375 )
376 .await
377 .expect("prefix-first reader must not wait for the missing v2 header bytes")
378 .unwrap_err();
379 assert!(matches!(
380 err,
381 FrameIoError::DecodeHeader(DecodeError::UnsupportedVersion { ver: 1 })
382 ));
383 }
384
385 #[tokio::test]
386 async fn invalid_header_is_typed_decode_error() {
387 let (mut client, mut server) = duplex(64);
388 let mut header = [0u8; HEADER_LEN];
389 header[4] = PROTOCOL_VERSION;
390 header[5] = 99;
391
392 let writer = tokio::spawn(async move {
393 client.write_all(&header).await.unwrap();
394 });
395
396 let err = read_frame(&mut server).await.unwrap_err();
397 writer.await.unwrap();
398 assert!(matches!(
399 err,
400 FrameIoError::DecodeHeader(DecodeError::UnknownFrameType { byte: 99 })
401 ));
402 }
403
404 #[tokio::test]
405 async fn eof_mid_body_is_typed_error() {
406 let (mut client, mut server) = duplex(64);
407 let frame = test_frame(1, 1, b"abcd");
408 let header = frame.header.encode();
409
410 let writer = tokio::spawn(async move {
411 client.write_all(&header).await.unwrap();
412 client.write_all(b"ab").await.unwrap();
413 });
414
415 let err = read_frame(&mut server).await.unwrap_err();
416 writer.await.unwrap();
417 assert!(matches!(
418 err,
419 FrameIoError::UnexpectedEof {
420 stage: ReadStage::Body,
421 expected: 4,
422 actual: 2
423 }
424 ));
425 }
426
427 #[tokio::test]
428 async fn pure_header_frame_with_body_len_is_typed_decode_error() {
429 let (mut client, mut server) = duplex(64);
430 let mut header = [0u8; HEADER_LEN];
431 header[0..4].copy_from_slice(&1u32.to_le_bytes());
432 header[4] = PROTOCOL_VERSION;
433 header[5] = FrameType::Ping as u8;
434 header[6] = Flags::new(false, Priority::Passive, false).0;
435
436 let writer = tokio::spawn(async move {
437 client.write_all(&header).await.unwrap();
438 });
439
440 let err = read_frame(&mut server).await.unwrap_err();
441 writer.await.unwrap();
442 assert!(matches!(
443 err,
444 FrameIoError::DecodeHeader(DecodeError::PureHeaderFrameWithBody {
445 ty: FrameType::Ping,
446 len: 1
447 })
448 ));
449 }
450
451 #[tokio::test]
452 async fn body_len_over_cap_is_rejected_before_allocation() {
453 let (mut client, mut server) = duplex(64);
454 let mut header = [0u8; HEADER_LEN];
455 header[0..4].copy_from_slice(&(MAX_FRAME_BODY_LEN + 1).to_le_bytes());
456 header[4] = PROTOCOL_VERSION;
457 header[5] = FrameType::Request as u8;
458 header[6] = Flags::new(false, Priority::Passive, false).0;
459
460 let writer = tokio::spawn(async move {
461 client.write_all(&header).await.unwrap();
462 });
463
464 let err = read_frame(&mut server).await.unwrap_err();
465 writer.await.unwrap();
466 assert!(matches!(
467 err,
468 FrameIoError::BodyTooLarge {
469 len,
470 max: MAX_FRAME_BODY_LEN
471 } if len == MAX_FRAME_BODY_LEN + 1
472 ));
473 }
474}