1use bytes::BytesMut;
21use nwd1::{decode, encode, Frame, MAGIC};
22use std::io;
23use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
24
25pub const MAX_FRAME_LEN_HARD: usize = 8 * 1024 * 1024;
27pub const DEFAULT_FRAME_LEN_SOFT: usize = 256 * 1024;
29const HEADER_LEN: usize = 8; pub async fn send_frame<W: AsyncWrite + Unpin>(writer: &mut W, frame: &Frame) -> io::Result<()> {
33 let data = encode(frame);
34 writer.write_all(&data).await?;
35 Ok(())
36}
37
38pub async fn recv_frame<R: AsyncRead + Unpin>(
41 reader: &mut R,
42 soft_cap: usize,
43) -> io::Result<Option<Frame>> {
44 let mut header = [0u8; HEADER_LEN];
46 if let Err(e) = reader.read_exact(&mut header).await {
47 return if e.kind() == io::ErrorKind::UnexpectedEof {
48 Ok(None)
49 } else {
50 Err(e)
51 };
52 }
53
54 if &header[..4] != MAGIC {
56 return Err(io::Error::new(io::ErrorKind::InvalidData, "nwd1 bad magic"));
57 }
58
59 let len = u32::from_be_bytes([header[4], header[5], header[6], header[7]]) as usize;
61
62 if len > soft_cap || len > MAX_FRAME_LEN_HARD {
64 return Err(io::Error::new(
65 io::ErrorKind::InvalidData,
66 "nwd1 frame too large",
67 ));
68 }
69
70 let mut body = vec![0u8; len];
72 reader.read_exact(&mut body).await?;
73
74 let mut buf = BytesMut::with_capacity(HEADER_LEN + len);
76 buf.extend_from_slice(&header);
77 buf.extend_from_slice(&body);
78
79 match decode(&buf.freeze()) {
80 Ok(frame) => Ok(Some(frame)),
81 Err(_) => Err(io::Error::new(
82 io::ErrorKind::InvalidData,
83 "nwd1 decode error",
84 )),
85 }
86}
87
88#[cfg(test)]
89mod tests {
90 use super::*;
91 use bytes::Bytes;
92 use netid64::NetId64;
93 #[tokio::test]
96 async fn roundtrip_via_inmemory() {
97 let frame = Frame {
99 id: NetId64::make(1, 7, 42),
100 kind: 1,
101 ver: 1,
102 payload: Bytes::from_static(b"hello"),
103 };
104
105 let data = encode(&frame);
107 let cursor = tokio::io::duplex(64 * 1024);
108 let (mut w, mut r) = cursor;
109
110 let written = data.clone();
112 let send = async move { w.write_all(&written).await };
113
114 let recv = async move { recv_frame(&mut r, DEFAULT_FRAME_LEN_SOFT).await };
116
117 let (sw, rr) = tokio::join!(send, recv);
118 sw.expect("write ok");
119 let decoded = rr.expect("io ok").expect("some");
120
121 assert_eq!(decoded.id.raw(), frame.id.raw());
122 assert_eq!(decoded.kind, frame.kind);
123 assert_eq!(decoded.ver, frame.ver);
124 assert_eq!(decoded.payload, frame.payload);
125 }
126}