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    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
71/// Writes one bounded length-prefixed JSON value.
72pub 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
163/// Rejects frames from incompatible clients before interpreting their message.
164pub 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}