mobius_gateway/wire/
codec.rs1use super::*;
2
3pub struct FrameReader<R> {
5 reader: R,
6 buffer: Vec<u8>,
7}
8
9impl<R> FrameReader<R> {
10 pub const fn new(reader: R) -> Self {
12 Self {
13 reader,
14 buffer: Vec::new(),
15 }
16 }
17}
18
19pub(super) fn deserialize_frame<'de, D>(
20 deserializer: D,
21) -> std::result::Result<(u16, Value), D::Error>
22where
23 D: serde::Deserializer<'de>,
24{
25 let Value::Object(mut object) = Value::deserialize(deserializer)? else {
26 return Err(D::Error::custom("gateway frame must be a JSON object"));
27 };
28 let version = object
29 .remove("version")
30 .ok_or_else(|| D::Error::missing_field("version"))?;
31 let version = serde_json::from_value(version).map_err(D::Error::custom)?;
32 Ok((version, Value::Object(object)))
33}
34
35pub async fn read_frame<T>(reader: &mut FrameReader<impl AsyncRead + Unpin>) -> Result<Option<T>>
37where
38 T: DeserializeOwned,
39{
40 read_frame_with_limit(reader, MAX_FRAME_BYTES).await
41}
42
43pub(crate) async fn read_frame_with_limit<T>(
44 reader: &mut FrameReader<impl AsyncRead + Unpin>,
45 max_bytes: usize,
46) -> Result<Option<T>>
47where
48 T: DeserializeOwned,
49{
50 loop {
51 let needed = if reader.buffer.len() >= 4 {
52 let prefix = reader.buffer[..4]
53 .try_into()
54 .map_err(|_| Error::Protocol("frame length is unsupported".into()))?;
55 let length = usize::try_from(u32::from_be_bytes(prefix))
56 .map_err(|_| Error::Protocol("frame length is unsupported".into()))?;
57 if length == 0 || length > max_bytes {
58 return Err(Error::Protocol(format!(
59 "frame length must be 1–{max_bytes} bytes"
60 )));
61 }
62 let frame_end = 4 + length;
63 if reader.buffer.len() >= frame_end {
64 let frame = serde_json::from_slice(&reader.buffer[4..frame_end])?;
65 reader.buffer.drain(..frame_end);
66 return Ok(Some(frame));
67 }
68 frame_end - reader.buffer.len()
69 } else {
70 4 - reader.buffer.len()
71 };
72 let mut chunk = [0_u8; 8 * 1024];
73 let chunk_bytes = needed.min(chunk.len());
74 let read = reader.reader.read(&mut chunk[..chunk_bytes]).await?;
75 if read == 0 {
76 if reader.buffer.is_empty() {
77 return Ok(None);
78 }
79 return Err(std::io::Error::from(std::io::ErrorKind::UnexpectedEof).into());
80 }
81 reader.buffer.extend_from_slice(&chunk[..read]);
82 }
83}
84
85pub async fn write_frame<T>(writer: &mut (impl AsyncWrite + Unpin), value: &T) -> Result<()>
87where
88 T: Serialize,
89{
90 let payload = serde_json::to_vec(value)?;
91 if payload.is_empty() || payload.len() > MAX_FRAME_BYTES {
92 return Err(Error::Protocol(format!(
93 "encoded frame must be 1–{MAX_FRAME_BYTES} bytes"
94 )));
95 }
96 let length = u32::try_from(payload.len())
97 .map_err(|_| Error::Protocol("encoded frame length is unsupported".into()))?;
98 writer.write_all(&length.to_be_bytes()).await?;
99 writer.write_all(&payload).await?;
100 writer.flush().await?;
101 Ok(())
102}
103
104pub(crate) async fn websocket_to_framed(
105 mut incoming: impl Stream<Item = std::result::Result<Message, WebSocketError>> + Unpin,
106 mut writer: impl AsyncWrite + Unpin,
107) -> Result<()> {
108 while let Some(message) = incoming.next().await {
109 match message.map_err(websocket_error)? {
110 Message::Binary(payload) if (1..=MAX_FRAME_BYTES).contains(&payload.len()) => {
111 let length = u32::try_from(payload.len())
112 .map_err(|_| Error::Protocol("WebSocket message is too large".into()))?;
113 writer.write_all(&length.to_be_bytes()).await?;
114 writer.write_all(&payload).await?;
115 }
116 Message::Ping(_) | Message::Pong(_) => {}
117 Message::Close(_) => return Ok(()),
118 Message::Binary(payload) => {
119 return Err(Error::Protocol(format!(
120 "WebSocket message length must be 1–{MAX_FRAME_BYTES} bytes, got {}",
121 payload.len()
122 )));
123 }
124 Message::Text(_) | Message::Frame(_) => {
125 return Err(Error::Protocol(
126 "WebSocket messages must be binary JSON frames".into(),
127 ));
128 }
129 }
130 }
131 Ok(())
132}
133
134pub(crate) async fn framed_to_websocket(
135 mut reader: impl AsyncRead + Unpin,
136 mut outgoing: impl Sink<Message, Error = WebSocketError> + Unpin,
137) -> Result<()> {
138 loop {
139 let mut prefix = [0_u8; 4];
140 let first =
141 match tokio::time::timeout(WEBSOCKET_KEEPALIVE_INTERVAL, reader.read(&mut prefix[..1]))
142 .await
143 {
144 Ok(read) => read?,
145 Err(_) => {
146 outgoing
147 .send(Message::Ping(Vec::new().into()))
148 .await
149 .map_err(websocket_error)?;
150 continue;
151 }
152 };
153 if first == 0 {
154 return outgoing.close().await.map_err(websocket_error);
155 }
156 reader.read_exact(&mut prefix[1..]).await?;
157 let length = usize::try_from(u32::from_be_bytes(prefix))
158 .map_err(|_| Error::Protocol("frame length is unsupported".into()))?;
159 if length == 0 || length > MAX_FRAME_BYTES {
160 return Err(Error::Protocol(format!(
161 "frame length must be 1–{MAX_FRAME_BYTES} bytes"
162 )));
163 }
164 let mut payload = vec![0_u8; length];
165 reader.read_exact(&mut payload).await?;
166 outgoing
167 .send(Message::Binary(payload.into()))
168 .await
169 .map_err(websocket_error)?;
170 }
171}
172
173pub(crate) fn websocket_error(error: WebSocketError) -> Error {
174 Error::Protocol(format!("WebSocket transport failed: {error}"))
175}
176
177pub fn validate_version(version: u16) -> Result<()> {
179 if version != PROTOCOL_VERSION {
180 return Err(Error::Protocol(format!(
181 "unsupported protocol version {version}; expected {PROTOCOL_VERSION}"
182 )));
183 }
184 Ok(())
185}
186
187pub(crate) fn validate_session_id(session_id: &str) -> Result<()> {
188 if session_id.trim().is_empty() || session_id.len() > 4 * 1024 {
189 return Err(Error::Config("session ID must be 1–4096 bytes".into()));
190 }
191 Ok(())
192}