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>
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 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 #[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 #[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 let mut expected = frame.header.encode().to_vec();
354 expected.extend_from_slice(&frame.body);
355 assert_eq!(writer.bytes, expected);
356 }
357
358 #[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}