1use serde::de::DeserializeOwned;
18use serde::Serialize;
19use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
20
21use crate::error::{DigPeerError, Result};
22
23pub const MAX_BODY: usize = 64 * 1024;
26
27pub async fn write_framed<W: AsyncWrite + Unpin>(w: &mut W, body: &[u8]) -> Result<()> {
29 if body.len() > MAX_BODY {
30 return Err(DigPeerError::Codec(format!(
31 "outbound body {} exceeds the {MAX_BODY}-byte bound",
32 body.len()
33 )));
34 }
35 w.write_all(&(body.len() as u32).to_be_bytes()).await?;
36 w.write_all(body).await?;
37 w.flush().await?;
38 Ok(())
39}
40
41pub async fn read_framed<R: AsyncRead + Unpin>(r: &mut R) -> Result<Vec<u8>> {
43 let mut len_buf = [0u8; 4];
44 r.read_exact(&mut len_buf).await?;
45 let len = u32::from_be_bytes(len_buf) as usize;
46 if len > MAX_BODY {
47 return Err(DigPeerError::Codec(format!(
48 "inbound body length {len} exceeds the {MAX_BODY}-byte bound"
49 )));
50 }
51 let mut body = vec![0u8; len];
52 r.read_exact(&mut body).await?;
53 Ok(body)
54}
55
56pub fn to_json<T: Serialize>(value: &T) -> Result<Vec<u8>> {
58 serde_json::to_vec(value).map_err(|e| DigPeerError::Codec(e.to_string()))
59}
60
61pub fn from_json<T: DeserializeOwned>(bytes: &[u8]) -> Result<T> {
63 serde_json::from_slice(bytes).map_err(|e| DigPeerError::Codec(e.to_string()))
64}
65
66#[cfg(test)]
67mod tests {
68 use super::*;
69
70 #[tokio::test]
73 async fn framed_body_round_trips() {
74 let body = b"{\"jsonrpc\":\"2.0\"}".to_vec();
75 let mut buf = Vec::new();
76 write_framed(&mut buf, &body).await.expect("write");
77 let mut cursor = std::io::Cursor::new(buf);
78 let read = read_framed(&mut cursor).await.expect("read");
79 assert_eq!(read, body);
80 }
81
82 #[tokio::test]
84 async fn empty_body_round_trips() {
85 let mut buf = Vec::new();
86 write_framed(&mut buf, &[]).await.expect("write");
87 let mut cursor = std::io::Cursor::new(buf);
88 assert!(read_framed(&mut cursor).await.expect("read").is_empty());
89 }
90
91 #[tokio::test]
94 async fn oversized_write_is_refused() {
95 let big = vec![0u8; MAX_BODY + 1];
96 let mut buf = Vec::new();
97 let result = write_framed(&mut buf, &big).await;
98 assert!(matches!(result, Err(DigPeerError::Codec(_))));
99 }
100
101 #[tokio::test]
104 async fn oversized_length_prefix_is_rejected() {
105 let mut framed = ((MAX_BODY + 1) as u32).to_be_bytes().to_vec();
106 framed.extend_from_slice(&[0u8; 8]);
107 let mut cursor = std::io::Cursor::new(framed);
108 assert!(matches!(
109 read_framed(&mut cursor).await,
110 Err(DigPeerError::Codec(_))
111 ));
112 }
113
114 #[test]
116 fn json_round_trips_and_rejects_garbage() {
117 let value = serde_json::json!({"a": 1});
118 let bytes = to_json(&value).expect("to_json");
119 let back: serde_json::Value = from_json(&bytes).expect("from_json");
120 assert_eq!(value, back);
121 assert!(matches!(
122 from_json::<serde_json::Value>(b"not json"),
123 Err(DigPeerError::Codec(_))
124 ));
125 }
126}