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. Header decode
93/// rules and the body-size cap are checked before emitting bytes. This function
94/// does not flush buffered writers; callers choose their own flush cadence.
95///
96/// HEADER AND BODY GO OUT AS ONE WRITE. Writing them separately looks harmless
97/// behind a `BufWriter` and is not: `BufWriter` passes any write at or above its
98/// capacity straight through to the socket, and flushes what it holds first to
99/// preserve ordering. A body larger than the buffer therefore emits the 21-byte
100/// header as a segment of its own, followed by the body as a second segment --
101/// the small-leading-segment shape that Nagle holds until an ACK returns. The
102/// boundary sits at the buffer capacity, so the same code path is fast for small
103/// frames and slow for large ones, which is the hardest version to notice.
104///
105/// Joining them also halves the syscalls on the unbuffered path, where every
106/// `write_all` is a syscall of its own.
107pub async fn write_frame<W>(writer: &mut W, frame: &Frame) -> Result<(), FrameIoError>
108where
109    W: AsyncWrite + Unpin,
110{
111    if frame.header.len as usize != frame.body.len() {
112        return Err(FrameIoError::BodyLengthMismatch {
113            header_len: frame.header.len,
114            body_len: frame.body.len(),
115        });
116    }
117
118    if frame.header.len > MAX_FRAME_BODY_LEN {
119        return Err(FrameIoError::BodyTooLarge {
120            len: frame.header.len,
121            max: MAX_FRAME_BODY_LEN,
122        });
123    }
124    let header = frame.header.encode();
125    // Frame fields are public, so callers can bypass the constructor. Validate
126    // before the first write to avoid poisoning an otherwise healthy stream.
127    decode_header(&header).map_err(FrameIoError::DecodeHeader)?;
128    if frame.body.is_empty() {
129        return writer.write_all(&header).await.map_err(FrameIoError::Io);
130    }
131
132    let mut joined = Vec::with_capacity(header.len() + frame.body.len());
133    joined.extend_from_slice(&header);
134    joined.extend_from_slice(&frame.body);
135    writer.write_all(&joined).await.map_err(FrameIoError::Io)
136}
137
138async fn read_exact_or_clean_eof<R>(
139    reader: &mut R,
140    buf: &mut [u8],
141    stage: ReadStage,
142) -> Result<bool, FrameIoError>
143where
144    R: AsyncRead + Unpin,
145{
146    let mut actual = 0;
147    while actual < buf.len() {
148        let n = reader
149            .read(&mut buf[actual..])
150            .await
151            .map_err(FrameIoError::Io)?;
152        if n == 0 {
153            if actual == 0 {
154                return Ok(false);
155            }
156            return Err(FrameIoError::UnexpectedEof {
157                stage,
158                expected: buf.len(),
159                actual,
160            });
161        }
162        actual += n;
163    }
164    Ok(true)
165}
166
167async fn read_exact_or_unexpected_eof<R>(
168    reader: &mut R,
169    buf: &mut [u8],
170    stage: ReadStage,
171) -> Result<(), FrameIoError>
172where
173    R: AsyncRead + Unpin,
174{
175    let mut actual = 0;
176    while actual < buf.len() {
177        let n = reader
178            .read(&mut buf[actual..])
179            .await
180            .map_err(FrameIoError::Io)?;
181        if n == 0 {
182            return Err(FrameIoError::UnexpectedEof {
183                stage,
184                expected: buf.len(),
185                actual,
186            });
187        }
188        actual += n;
189    }
190    Ok(())
191}
192
193impl fmt::Display for FrameIoError {
194    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
195        match self {
196            Self::Io(err) => write!(f, "frame I/O error: {err}"),
197            Self::DecodeHeader(err) => write!(f, "invalid envelope header: {err}"),
198            Self::BodyTooLarge { len, max } => {
199                write!(f, "frame body length {len} exceeds max {max}")
200            }
201            Self::UnexpectedEof {
202                stage,
203                expected,
204                actual,
205            } => write!(
206                f,
207                "unexpected EOF while reading {stage:?}: expected {expected} bytes, got {actual}"
208            ),
209            Self::BodyLengthMismatch {
210                header_len,
211                body_len,
212            } => write!(
213                f,
214                "frame header len ({header_len}) does not match body length ({body_len})"
215            ),
216        }
217    }
218}
219
220impl Error for FrameIoError {
221    fn source(&self) -> Option<&(dyn Error + 'static)> {
222        match self {
223            Self::Io(err) => Some(err),
224            Self::DecodeHeader(err) => Some(err),
225            Self::UnexpectedEof { .. }
226            | Self::BodyTooLarge { .. }
227            | Self::BodyLengthMismatch { .. } => None,
228        }
229    }
230}
231
232impl From<io::Error> for FrameIoError {
233    fn from(err: io::Error) -> Self {
234        Self::Io(err)
235    }
236}
237
238#[cfg(test)]
239mod tests {
240    use super::*;
241    use subc_protocol::{Flags, FrameType, Priority, PROTOCOL_VERSION};
242    use tokio::io::{duplex, AsyncWriteExt};
243
244    fn test_frame(channel: u16, corr: u64, body: &[u8]) -> Frame {
245        Frame::build(
246            FrameType::Request,
247            Flags::new(true, Priority::Interactive, false),
248            channel,
249            1,
250            corr,
251            body.to_vec(),
252        )
253        .unwrap()
254    }
255
256    #[tokio::test]
257    async fn writer_refuses_invalid_public_frames_before_emitting_bytes() {
258        let valid = test_frame(1, 1, b"x");
259        let mut pure = valid.clone();
260        pure.header.ty = FrameType::Ping;
261        let mut control = valid.clone();
262        control.header.channel = 0;
263        let mut sheddable = valid.clone();
264        sheddable.header.flags = sheddable
265            .header
266            .flags
267            .with_admission_class(subc_protocol::AdmissionClass::Sheddable);
268        for invalid in [pure, control, sheddable] {
269            let mut writer = WriteCounter::default();
270            assert!(
271                write_frame(&mut writer, &invalid).await.is_err(),
272                "{:?}",
273                invalid.header
274            );
275            assert!(writer.bytes.is_empty());
276        }
277    }
278
279    #[tokio::test]
280    async fn writer_refuses_oversized_public_frames_before_emitting_bytes() {
281        let mut frame = test_frame(1, 1, b"");
282        frame.body = vec![0; MAX_FRAME_BODY_LEN as usize + 1];
283        frame.header.len = MAX_FRAME_BODY_LEN + 1;
284        let mut writer = WriteCounter::default();
285        assert!(matches!(
286            write_frame(&mut writer, &frame).await,
287            Err(FrameIoError::BodyTooLarge { .. })
288        ));
289        assert!(writer.bytes.is_empty());
290    }
291
292    /// Counts `poll_write` calls and records what each one carried, which is the
293    /// only way to observe segmentation: every round-trip test passes whether a
294    /// frame goes out as one write or as twenty, because the reader reassembles
295    /// either way. The bytes are identical and the latency is not.
296    #[derive(Default)]
297    struct WriteCounter {
298        writes: Vec<usize>,
299        bytes: Vec<u8>,
300    }
301
302    impl AsyncWrite for WriteCounter {
303        fn poll_write(
304            mut self: std::pin::Pin<&mut Self>,
305            _cx: &mut std::task::Context<'_>,
306            buf: &[u8],
307        ) -> std::task::Poll<io::Result<usize>> {
308            self.writes.push(buf.len());
309            self.bytes.extend_from_slice(buf);
310            std::task::Poll::Ready(Ok(buf.len()))
311        }
312
313        fn poll_flush(
314            self: std::pin::Pin<&mut Self>,
315            _cx: &mut std::task::Context<'_>,
316        ) -> std::task::Poll<io::Result<()>> {
317            std::task::Poll::Ready(Ok(()))
318        }
319
320        fn poll_shutdown(
321            self: std::pin::Pin<&mut Self>,
322            _cx: &mut std::task::Context<'_>,
323        ) -> std::task::Poll<io::Result<()>> {
324            std::task::Poll::Ready(Ok(()))
325        }
326    }
327
328    /// A frame with a body must reach the socket as ONE write.
329    ///
330    /// Writing the header separately is correct and slow: behind a `BufWriter` a
331    /// body at or above the buffer capacity is passed straight through, and the
332    /// buffered header is flushed first to keep ordering -- so the header goes out
333    /// alone as a 21-byte segment, and Nagle holds the body until that segment is
334    /// acknowledged. The reader cannot tell the difference, so nothing else in the
335    /// suite can fail when this regresses.
336    #[tokio::test]
337    async fn a_frame_with_a_body_reaches_the_socket_as_one_write() {
338        let mut writer = WriteCounter::default();
339        let frame = test_frame(3, 11, &vec![0xABu8; 16 * 1024]);
340
341        write_frame(&mut writer, &frame).await.unwrap();
342
343        assert_eq!(
344            writer.writes.len(),
345            1,
346            "header and body must be one write, got segments {:?}",
347            writer.writes
348        );
349        assert_eq!(writer.writes[0], HEADER_LEN + frame.body.len());
350
351        // The joined buffer must still be the header followed by the body, or the
352        // single-write assertion above would be satisfied by writing anything once.
353        let mut expected = frame.header.encode().to_vec();
354        expected.extend_from_slice(&frame.body);
355        assert_eq!(writer.bytes, expected);
356    }
357
358    /// A bodyless frame writes only the header, and must not gain a second empty
359    /// write from the joining path.
360    #[tokio::test]
361    async fn a_bodyless_frame_writes_only_its_header() {
362        let mut writer = WriteCounter::default();
363        let frame = test_frame(4, 12, b"");
364
365        write_frame(&mut writer, &frame).await.unwrap();
366
367        assert_eq!(writer.writes, vec![HEADER_LEN]);
368    }
369
370    #[tokio::test]
371    async fn read_write_round_trip_preserves_opaque_body() {
372        let (mut client, mut server) = duplex(128);
373        let frame = test_frame(7, 42, b"opaque\0json? no parse");
374        let expected = frame.clone();
375
376        let writer = tokio::spawn(async move { write_frame(&mut client, &frame).await });
377        let read = read_frame(&mut server).await.unwrap().unwrap();
378
379        writer.await.unwrap().unwrap();
380        assert_eq!(read, expected);
381    }
382
383    #[tokio::test]
384    async fn partial_header_and_body_are_assembled() {
385        let (mut client, mut server) = duplex(128);
386        let frame = test_frame(2, 99, b"chunked-body");
387        let mut bytes = frame.header.encode().to_vec();
388        bytes.extend_from_slice(&frame.body);
389        let expected = frame.clone();
390
391        let writer = tokio::spawn(async move {
392            client.write_all(&bytes[..3]).await.unwrap();
393            client.write_all(&bytes[3..10]).await.unwrap();
394            client.write_all(&bytes[10..]).await.unwrap();
395        });
396
397        let read = read_frame(&mut server).await.unwrap().unwrap();
398        writer.await.unwrap();
399        assert_eq!(read, expected);
400    }
401
402    #[tokio::test]
403    async fn clean_eof_before_header_returns_none() {
404        let (client, mut server) = duplex(16);
405        drop(client);
406
407        assert!(read_frame(&mut server).await.unwrap().is_none());
408    }
409
410    #[tokio::test]
411    async fn stale_v1_pure_header_is_rejected_from_prefix_without_waiting() {
412        let (mut client, mut server) = duplex(64);
413        let mut stale_header = [0u8; 17];
414        stale_header[4] = 1;
415        stale_header[5] = FrameType::Ping as u8;
416        client.write_all(&stale_header).await.unwrap();
417
418        let err = tokio::time::timeout(
419            std::time::Duration::from_millis(100),
420            read_frame(&mut server),
421        )
422        .await
423        .expect("prefix-first reader must not wait for the missing v2 header bytes")
424        .unwrap_err();
425        assert!(matches!(
426            err,
427            FrameIoError::DecodeHeader(DecodeError::UnsupportedVersion { ver: 1 })
428        ));
429    }
430
431    #[tokio::test]
432    async fn invalid_header_is_typed_decode_error() {
433        let (mut client, mut server) = duplex(64);
434        let mut header = [0u8; HEADER_LEN];
435        header[4] = PROTOCOL_VERSION;
436        header[5] = 99;
437
438        let writer = tokio::spawn(async move {
439            client.write_all(&header).await.unwrap();
440        });
441
442        let err = read_frame(&mut server).await.unwrap_err();
443        writer.await.unwrap();
444        assert!(matches!(
445            err,
446            FrameIoError::DecodeHeader(DecodeError::UnknownFrameType { byte: 99 })
447        ));
448    }
449
450    #[tokio::test]
451    async fn eof_mid_body_is_typed_error() {
452        let (mut client, mut server) = duplex(64);
453        let frame = test_frame(1, 1, b"abcd");
454        let header = frame.header.encode();
455
456        let writer = tokio::spawn(async move {
457            client.write_all(&header).await.unwrap();
458            client.write_all(b"ab").await.unwrap();
459        });
460
461        let err = read_frame(&mut server).await.unwrap_err();
462        writer.await.unwrap();
463        assert!(matches!(
464            err,
465            FrameIoError::UnexpectedEof {
466                stage: ReadStage::Body,
467                expected: 4,
468                actual: 2
469            }
470        ));
471    }
472
473    #[tokio::test]
474    async fn pure_header_frame_with_body_len_is_typed_decode_error() {
475        let (mut client, mut server) = duplex(64);
476        let mut header = [0u8; HEADER_LEN];
477        header[0..4].copy_from_slice(&1u32.to_le_bytes());
478        header[4] = PROTOCOL_VERSION;
479        header[5] = FrameType::Ping as u8;
480        header[6] = Flags::new(false, Priority::Passive, false).0;
481
482        let writer = tokio::spawn(async move {
483            client.write_all(&header).await.unwrap();
484        });
485
486        let err = read_frame(&mut server).await.unwrap_err();
487        writer.await.unwrap();
488        assert!(matches!(
489            err,
490            FrameIoError::DecodeHeader(DecodeError::PureHeaderFrameWithBody {
491                ty: FrameType::Ping,
492                len: 1
493            })
494        ));
495    }
496
497    #[tokio::test]
498    async fn body_len_over_cap_is_rejected_before_allocation() {
499        let (mut client, mut server) = duplex(64);
500        let mut header = [0u8; HEADER_LEN];
501        header[0..4].copy_from_slice(&(MAX_FRAME_BODY_LEN + 1).to_le_bytes());
502        header[4] = PROTOCOL_VERSION;
503        header[5] = FrameType::Request as u8;
504        header[6] = Flags::new(false, Priority::Passive, false).0;
505
506        let writer = tokio::spawn(async move {
507            client.write_all(&header).await.unwrap();
508        });
509
510        let err = read_frame(&mut server).await.unwrap_err();
511        writer.await.unwrap();
512        assert!(matches!(
513            err,
514            FrameIoError::BodyTooLarge {
515                len,
516                max: MAX_FRAME_BODY_LEN
517            } if len == MAX_FRAME_BODY_LEN + 1
518        ));
519    }
520}