Skip to main content

subc_transport/
frame_io.rs

1//! Async envelope frame I/O over the authenticated stream.
2//!
3//! `read_frame`/`write_frame` are the post-handshake continuation of
4//! [`authenticate_client`](crate::authenticate_client)/`authenticate_server` on
5//! the same socket: once the connection is authenticated, both peers exchange
6//! [`Frame`]s (the 21-byte envelope header + opaque body). The framing codec is
7//! shared by subc-core and modules (AFT) so the wire cannot drift.
8
9use 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/// Which part of a frame was being read when EOF arrived.
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
19pub enum ReadStage {
20    Header,
21    Body,
22}
23
24/// Errors from async envelope frame I/O.
25#[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
44/// Read one complete frame from an async stream.
45///
46/// Returns `Ok(None)` only for a clean EOF before the next header begins. EOF
47/// after any header byte, or before all body bytes arrive, is a typed
48/// [`FrameIoError::UnexpectedEof`]. The body is returned as opaque bytes.
49pub 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
89/// Write one complete frame to an async stream.
90///
91/// The header's `len` must match the opaque body length; mismatches are reported
92/// as a typed error rather than silently rewriting the header. This function does
93/// not flush buffered writers; callers choose their own flush cadence.
94///
95/// HEADER AND BODY GO OUT AS ONE WRITE. Writing them separately looks harmless
96/// behind a `BufWriter` and is not: `BufWriter` passes any write at or above its
97/// capacity straight through to the socket, and flushes what it holds first to
98/// preserve ordering. A body larger than the buffer therefore emits the 21-byte
99/// header as a segment of its own, followed by the body as a second segment --
100/// the small-leading-segment shape that Nagle holds until an ACK returns. The
101/// boundary sits at the buffer capacity, so the same code path is fast for small
102/// frames and slow for large ones, which is the hardest version to notice.
103///
104/// Joining them also halves the syscalls on the unbuffered path, where every
105/// `write_all` is a syscall of its own.
106pub 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    /// Counts `poll_write` calls and records what each one carried, which is the
247    /// only way to observe segmentation: every round-trip test passes whether a
248    /// frame goes out as one write or as twenty, because the reader reassembles
249    /// either way. The bytes are identical and the latency is not.
250    #[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    /// A frame with a body must reach the socket as ONE write.
283    ///
284    /// Writing the header separately is correct and slow: behind a `BufWriter` a
285    /// body at or above the buffer capacity is passed straight through, and the
286    /// buffered header is flushed first to keep ordering -- so the header goes out
287    /// alone as a 21-byte segment, and Nagle holds the body until that segment is
288    /// acknowledged. The reader cannot tell the difference, so nothing else in the
289    /// suite can fail when this regresses.
290    #[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        // The joined buffer must still be the header followed by the body, or the
306        // single-write assertion above would be satisfied by writing anything once.
307        let mut expected = frame.header.encode().to_vec();
308        expected.extend_from_slice(&frame.body);
309        assert_eq!(writer.bytes, expected);
310    }
311
312    /// A bodyless frame writes only the header, and must not gain a second empty
313    /// write from the joining path.
314    #[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}