Skip to main content

mobius_gateway/wire/
codec.rs

1use super::*;
2
3/// Cancellation-safe reader for length-prefixed gateway frames.
4pub struct FrameReader<R> {
5    reader: R,
6    buffer: Vec<u8>,
7}
8
9impl<R> FrameReader<R> {
10    /// Wraps one transport reader and retains partial frames between reads.
11    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
35/// Reads one length-prefixed JSON value, returning `None` only for a clean EOF.
36pub 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
85/// Writes one bounded length-prefixed JSON value.
86pub 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
177/// Rejects frames from incompatible clients before interpreting their message.
178pub 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}