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 loop {
41 if reader.buffer.len() >= 4 {
42 let prefix = reader.buffer[..4]
43 .try_into()
44 .map_err(|_| Error::Protocol("frame length is unsupported".into()))?;
45 let length = usize::try_from(u32::from_be_bytes(prefix))
46 .map_err(|_| Error::Protocol("frame length is unsupported".into()))?;
47 if length == 0 || length > MAX_FRAME_BYTES {
48 return Err(Error::Protocol(format!(
49 "frame length must be 1–{MAX_FRAME_BYTES} bytes"
50 )));
51 }
52 let frame_end = 4 + length;
53 if reader.buffer.len() >= frame_end {
54 let frame = serde_json::from_slice(&reader.buffer[4..frame_end])?;
55 reader.buffer.drain(..frame_end);
56 return Ok(Some(frame));
57 }
58 }
59 let mut chunk = [0_u8; 8 * 1024];
60 let read = reader.reader.read(&mut chunk).await?;
61 if read == 0 {
62 if reader.buffer.is_empty() {
63 return Ok(None);
64 }
65 return Err(std::io::Error::from(std::io::ErrorKind::UnexpectedEof).into());
66 }
67 reader.buffer.extend_from_slice(&chunk[..read]);
68 }
69}
70
71pub async fn write_frame<T>(writer: &mut (impl AsyncWrite + Unpin), value: &T) -> Result<()>
73where
74 T: Serialize,
75{
76 let payload = serde_json::to_vec(value)?;
77 if payload.is_empty() || payload.len() > MAX_FRAME_BYTES {
78 return Err(Error::Protocol(format!(
79 "encoded frame must be 1–{MAX_FRAME_BYTES} bytes"
80 )));
81 }
82 let length = u32::try_from(payload.len())
83 .map_err(|_| Error::Protocol("encoded frame length is unsupported".into()))?;
84 writer.write_all(&length.to_be_bytes()).await?;
85 writer.write_all(&payload).await?;
86 writer.flush().await?;
87 Ok(())
88}
89
90pub(crate) async fn websocket_to_framed(
91 mut incoming: impl Stream<Item = std::result::Result<Message, WebSocketError>> + Unpin,
92 mut writer: impl AsyncWrite + Unpin,
93) -> Result<()> {
94 while let Some(message) = incoming.next().await {
95 match message.map_err(websocket_error)? {
96 Message::Binary(payload) if (1..=MAX_FRAME_BYTES).contains(&payload.len()) => {
97 let length = u32::try_from(payload.len())
98 .map_err(|_| Error::Protocol("WebSocket message is too large".into()))?;
99 writer.write_all(&length.to_be_bytes()).await?;
100 writer.write_all(&payload).await?;
101 }
102 Message::Ping(_) | Message::Pong(_) => {}
103 Message::Close(_) => return Ok(()),
104 Message::Binary(payload) => {
105 return Err(Error::Protocol(format!(
106 "WebSocket message length must be 1–{MAX_FRAME_BYTES} bytes, got {}",
107 payload.len()
108 )));
109 }
110 Message::Text(_) | Message::Frame(_) => {
111 return Err(Error::Protocol(
112 "WebSocket messages must be binary JSON frames".into(),
113 ));
114 }
115 }
116 }
117 Ok(())
118}
119
120pub(crate) async fn framed_to_websocket(
121 mut reader: impl AsyncRead + Unpin,
122 mut outgoing: impl Sink<Message, Error = WebSocketError> + Unpin,
123) -> Result<()> {
124 loop {
125 let mut prefix = [0_u8; 4];
126 let first =
127 match tokio::time::timeout(WEBSOCKET_KEEPALIVE_INTERVAL, reader.read(&mut prefix[..1]))
128 .await
129 {
130 Ok(read) => read?,
131 Err(_) => {
132 outgoing
133 .send(Message::Ping(Vec::new().into()))
134 .await
135 .map_err(websocket_error)?;
136 continue;
137 }
138 };
139 if first == 0 {
140 return outgoing.close().await.map_err(websocket_error);
141 }
142 reader.read_exact(&mut prefix[1..]).await?;
143 let length = usize::try_from(u32::from_be_bytes(prefix))
144 .map_err(|_| Error::Protocol("frame length is unsupported".into()))?;
145 if length == 0 || length > MAX_FRAME_BYTES {
146 return Err(Error::Protocol(format!(
147 "frame length must be 1–{MAX_FRAME_BYTES} bytes"
148 )));
149 }
150 let mut payload = vec![0_u8; length];
151 reader.read_exact(&mut payload).await?;
152 outgoing
153 .send(Message::Binary(payload.into()))
154 .await
155 .map_err(websocket_error)?;
156 }
157}
158
159pub(crate) fn websocket_error(error: WebSocketError) -> Error {
160 Error::Protocol(format!("WebSocket transport failed: {error}"))
161}
162
163pub fn validate_version(version: u16) -> Result<()> {
165 if version != PROTOCOL_VERSION {
166 return Err(Error::Protocol(format!(
167 "unsupported protocol version {version}; expected {PROTOCOL_VERSION}"
168 )));
169 }
170 Ok(())
171}
172
173pub(crate) fn validate_session_id(session_id: &str) -> Result<()> {
174 if session_id.trim().is_empty() || session_id.len() > 4 * 1024 {
175 return Err(Error::Config("session ID must be 1–4096 bytes".into()));
176 }
177 Ok(())
178}