use serde::de::DeserializeOwned;
use serde::Serialize;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use crate::error::{DigPeerError, Result};
pub const MAX_BODY: usize = 64 * 1024;
pub async fn write_framed<W: AsyncWrite + Unpin>(w: &mut W, body: &[u8]) -> Result<()> {
if body.len() > MAX_BODY {
return Err(DigPeerError::Codec(format!(
"outbound body {} exceeds the {MAX_BODY}-byte bound",
body.len()
)));
}
w.write_all(&(body.len() as u32).to_be_bytes()).await?;
w.write_all(body).await?;
w.flush().await?;
Ok(())
}
pub async fn read_framed<R: AsyncRead + Unpin>(r: &mut R) -> Result<Vec<u8>> {
let mut len_buf = [0u8; 4];
r.read_exact(&mut len_buf).await?;
let len = u32::from_be_bytes(len_buf) as usize;
if len > MAX_BODY {
return Err(DigPeerError::Codec(format!(
"inbound body length {len} exceeds the {MAX_BODY}-byte bound"
)));
}
let mut body = vec![0u8; len];
r.read_exact(&mut body).await?;
Ok(body)
}
pub fn to_json<T: Serialize>(value: &T) -> Result<Vec<u8>> {
serde_json::to_vec(value).map_err(|e| DigPeerError::Codec(e.to_string()))
}
pub fn from_json<T: DeserializeOwned>(bytes: &[u8]) -> Result<T> {
serde_json::from_slice(bytes).map_err(|e| DigPeerError::Codec(e.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn framed_body_round_trips() {
let body = b"{\"jsonrpc\":\"2.0\"}".to_vec();
let mut buf = Vec::new();
write_framed(&mut buf, &body).await.expect("write");
let mut cursor = std::io::Cursor::new(buf);
let read = read_framed(&mut cursor).await.expect("read");
assert_eq!(read, body);
}
#[tokio::test]
async fn empty_body_round_trips() {
let mut buf = Vec::new();
write_framed(&mut buf, &[]).await.expect("write");
let mut cursor = std::io::Cursor::new(buf);
assert!(read_framed(&mut cursor).await.expect("read").is_empty());
}
#[tokio::test]
async fn oversized_write_is_refused() {
let big = vec![0u8; MAX_BODY + 1];
let mut buf = Vec::new();
let result = write_framed(&mut buf, &big).await;
assert!(matches!(result, Err(DigPeerError::Codec(_))));
}
#[tokio::test]
async fn oversized_length_prefix_is_rejected() {
let mut framed = ((MAX_BODY + 1) as u32).to_be_bytes().to_vec();
framed.extend_from_slice(&[0u8; 8]);
let mut cursor = std::io::Cursor::new(framed);
assert!(matches!(
read_framed(&mut cursor).await,
Err(DigPeerError::Codec(_))
));
}
#[test]
fn json_round_trips_and_rejects_garbage() {
let value = serde_json::json!({"a": 1});
let bytes = to_json(&value).expect("to_json");
let back: serde_json::Value = from_json(&bytes).expect("from_json");
assert_eq!(value, back);
assert!(matches!(
from_json::<serde_json::Value>(b"not json"),
Err(DigPeerError::Codec(_))
));
}
}