1use anyhow::{Context, Result, bail};
16use serde::de::DeserializeOwned;
17use serde::{Deserialize, Serialize};
18use yaiba_core::{Entry, VersionVector};
19
20const MAX_FRAME: usize = 32 * 1024 * 1024;
24
25#[derive(Debug, Serialize, Deserialize)]
26pub struct Hello {
27 pub room: String,
30 pub vv: VersionVector,
31}
32
33#[derive(Debug, Serialize, Deserialize)]
34pub struct Offer {
35 pub vv: VersionVector,
36 pub entries: Vec<Entry>,
37}
38
39#[derive(Debug, Serialize, Deserialize)]
40pub struct Push {
41 pub entries: Vec<Entry>,
42}
43
44pub async fn write_frame<T, W>(writer: &mut W, value: &T) -> Result<()>
48where
49 T: Serialize,
50 W: tokio::io::AsyncWrite + Unpin,
51{
52 use tokio::io::AsyncWriteExt;
53 let bytes = serde_json::to_vec(value).context("encode frame")?;
54 if bytes.len() > MAX_FRAME {
55 bail!("frame too large to send: {} bytes", bytes.len());
56 }
57 writer
58 .write_all(&(bytes.len() as u32).to_be_bytes())
59 .await?;
60 writer.write_all(&bytes).await?;
61 writer.flush().await?;
62 Ok(())
63}
64
65pub async fn read_frame<T, R>(reader: &mut R) -> Result<T>
66where
67 T: DeserializeOwned,
68 R: tokio::io::AsyncRead + Unpin,
69{
70 use tokio::io::AsyncReadExt;
71 let mut len = [0u8; 4];
72 reader.read_exact(&mut len).await.context("read length")?;
73 let len = u32::from_be_bytes(len) as usize;
74 if len > MAX_FRAME {
75 bail!("peer announced an oversized frame: {len} bytes");
76 }
77 let mut buf = vec![0u8; len];
78 reader.read_exact(&mut buf).await.context("read body")?;
79 serde_json::from_slice(&buf).context("decode frame")
80}
81
82pub fn room_matches(a: &str, b: &str) -> bool {
89 if a.len() != b.len() {
90 return false;
91 }
92 a.bytes()
93 .zip(b.bytes())
94 .fold(0u8, |acc, (x, y)| acc | (x ^ y))
95 == 0
96}
97
98#[cfg(test)]
99mod tests {
100 use super::*;
101
102 #[tokio::test]
103 async fn frames_round_trip() {
104 let hello = Hello {
105 room: "abc".into(),
106 vv: VersionVector::new(),
107 };
108 let mut buf = Vec::new();
109 write_frame(&mut buf, &hello).await.unwrap();
110
111 let mut cursor = std::io::Cursor::new(buf);
112 let back: Hello = read_frame(&mut cursor).await.unwrap();
113 assert_eq!(back.room, "abc");
114 }
115
116 #[tokio::test]
117 async fn an_oversized_length_prefix_is_refused_before_allocating() {
118 let mut framed = Vec::new();
119 framed.extend_from_slice(&u32::MAX.to_be_bytes());
120 let mut cursor = std::io::Cursor::new(framed);
121 let err = read_frame::<Hello, _>(&mut cursor).await.unwrap_err();
122 assert!(err.to_string().contains("oversized"));
123 }
124
125 #[test]
126 fn room_comparison_rejects_mismatches_and_length_changes() {
127 assert!(room_matches("deadbeef", "deadbeef"));
128 assert!(!room_matches("deadbeef", "deadbeee"));
129 assert!(!room_matches("deadbeef", "deadbee"));
130 }
131}