Skip to main content

mobius_gateway/wire/
codec.rs

1use super::*;
2
3/// An immutable session frame shared by replay and live subscribers.
4#[derive(Debug, Clone)]
5pub(crate) struct SharedFrame {
6    frame: std::sync::Arc<ServerFrame>,
7    payload: Option<std::sync::Arc<[u8]>>,
8}
9
10impl SharedFrame {
11    pub(crate) fn new(frame: ServerFrame) -> Self {
12        Self {
13            frame: std::sync::Arc::new(frame),
14            payload: None,
15        }
16    }
17
18    pub(crate) fn encoded(frame: ServerFrame) -> Result<Self> {
19        let payload = serde_json::to_vec(&frame)?;
20        if payload.len() > MAX_FRAME_BYTES {
21            return Err(Error::Protocol(format!(
22                "agent event exceeds the {MAX_FRAME_BYTES}-byte gateway frame limit"
23            )));
24        }
25        Ok(Self {
26            frame: std::sync::Arc::new(frame),
27            payload: Some(payload.into()),
28        })
29    }
30
31    pub(crate) fn encoded_len(&self) -> usize {
32        self.payload
33            .as_ref()
34            .expect("replay frames are encoded")
35            .len()
36    }
37
38    pub(crate) async fn write(&self, writer: &mut (impl AsyncWrite + Unpin)) -> Result<()> {
39        match &self.payload {
40            Some(payload) => write_payload(writer, payload).await,
41            None => write_frame(writer, &*self.frame).await,
42        }
43    }
44}
45
46impl std::ops::Deref for SharedFrame {
47    type Target = ServerFrame;
48
49    fn deref(&self) -> &Self::Target {
50        &self.frame
51    }
52}
53
54/// Cancellation-safe reader for length-prefixed gateway frames.
55pub struct FrameReader<R> {
56    reader: R,
57    buffer: Vec<u8>,
58}
59
60impl<R> FrameReader<R> {
61    /// Wraps one transport reader and retains partial frames between reads.
62    pub const fn new(reader: R) -> Self {
63        Self {
64            reader,
65            buffer: Vec::new(),
66        }
67    }
68}
69
70pub(super) fn deserialize_frame<'de, D>(
71    deserializer: D,
72) -> std::result::Result<(u16, Value), D::Error>
73where
74    D: serde::Deserializer<'de>,
75{
76    let Value::Object(mut object) = Value::deserialize(deserializer)? else {
77        return Err(D::Error::custom("gateway frame must be a JSON object"));
78    };
79    let version = object
80        .remove("version")
81        .ok_or_else(|| D::Error::missing_field("version"))?;
82    let version = serde_json::from_value(version).map_err(D::Error::custom)?;
83    Ok((version, Value::Object(object)))
84}
85
86/// Reads one length-prefixed JSON value, returning `None` only for a clean EOF.
87/// # Errors
88///
89/// Returns an error if the resource cannot be read, decoded, or validated.
90pub async fn read_frame<T>(reader: &mut FrameReader<impl AsyncRead + Unpin>) -> Result<Option<T>>
91where
92    T: DeserializeOwned,
93{
94    read_frame_with_limit(reader, MAX_FRAME_BYTES).await
95}
96
97pub(crate) async fn read_frame_with_limit<T>(
98    reader: &mut FrameReader<impl AsyncRead + Unpin>,
99    max_bytes: usize,
100) -> Result<Option<T>>
101where
102    T: DeserializeOwned,
103{
104    loop {
105        let needed = if reader.buffer.len() >= 4 {
106            let prefix = reader.buffer[..4]
107                .try_into()
108                .map_err(|_| Error::Protocol("frame length is unsupported".into()))?;
109            let length = usize::try_from(u32::from_be_bytes(prefix))
110                .map_err(|_| Error::Protocol("frame length is unsupported".into()))?;
111            if length == 0 || length > max_bytes {
112                return Err(Error::Protocol(format!(
113                    "frame length must be 1–{max_bytes} bytes"
114                )));
115            }
116            let frame_end = 4 + length;
117            if reader.buffer.len() >= frame_end {
118                let frame = serde_json::from_slice(&reader.buffer[4..frame_end])?;
119                reader.buffer.clear();
120                // Reuse ordinary frames without pinning bulk-transfer memory for idle clients.
121                if reader.buffer.capacity() > 64 * 1024 {
122                    reader.buffer = Vec::new();
123                }
124                return Ok(Some(frame));
125            }
126            frame_end - reader.buffer.len()
127        } else {
128            4 - reader.buffer.len()
129        };
130        let mut chunk = [0_u8; 8 * 1024];
131        let chunk_bytes = needed.min(chunk.len());
132        let read = reader.reader.read(&mut chunk[..chunk_bytes]).await?;
133        if read == 0 {
134            if reader.buffer.is_empty() {
135                return Ok(None);
136            }
137            return Err(std::io::Error::from(std::io::ErrorKind::UnexpectedEof).into());
138        }
139        reader.buffer.extend_from_slice(&chunk[..read]);
140    }
141}
142
143/// Writes one bounded length-prefixed JSON value.
144/// Discard the connection after a write error: the frame may be partially written.
145/// # Errors
146///
147/// Returns an error if the value cannot be encoded or persisted.
148pub async fn write_frame<T>(writer: &mut (impl AsyncWrite + Unpin), value: &T) -> Result<()>
149where
150    T: Serialize,
151{
152    let payload = serde_json::to_vec(value)?;
153    write_payload(writer, &payload).await
154}
155
156pub(crate) async fn write_payload(
157    writer: &mut (impl AsyncWrite + Unpin),
158    payload: &[u8],
159) -> Result<()> {
160    if payload.is_empty() || payload.len() > MAX_FRAME_BYTES {
161        return Err(Error::Protocol(format!(
162            "encoded frame must be 1–{MAX_FRAME_BYTES} bytes"
163        )));
164    }
165    let length = u32::try_from(payload.len())
166        .map_err(|_| Error::Protocol("encoded frame length is unsupported".into()))?;
167    tokio::time::timeout(WRITE_TIMEOUT, async {
168        writer.write_all(&length.to_be_bytes()).await?;
169        writer.write_all(payload).await?;
170        writer.flush().await
171    })
172    .await
173    .map_err(|_| std::io::Error::from(std::io::ErrorKind::TimedOut))??;
174    Ok(())
175}
176
177pub(crate) fn websocket_error(error: WebSocketError) -> Error {
178    let kind = match error {
179        WebSocketError::Io(error) => {
180            return Error::Protocol(format!("WebSocket I/O failure: {:?}", error.kind()));
181        }
182        WebSocketError::Http(response) => {
183            return Error::WebSocketUpgrade {
184                status: response.status().as_u16(),
185                retry_after: retry_after(response.headers()),
186            };
187        }
188        WebSocketError::ConnectionClosed | WebSocketError::AlreadyClosed => "closed",
189        WebSocketError::Tls(_) => "TLS",
190        WebSocketError::Capacity(_) => "capacity",
191        WebSocketError::Protocol(_) => "protocol",
192        WebSocketError::WriteBufferFull(_) => "write buffer",
193        WebSocketError::Utf8(_) => "UTF-8",
194        WebSocketError::AttackAttempt => "attack rejected",
195        WebSocketError::Url(_) => "URL",
196        WebSocketError::HttpFormat(_) => "HTTP format",
197    };
198    Error::Protocol(format!("WebSocket {kind} failure"))
199}
200
201fn retry_after(headers: &tokio_tungstenite::tungstenite::http::HeaderMap) -> Option<Duration> {
202    let values = headers.get_all("retry-after");
203    if values.iter().count() != 1 {
204        return None;
205    }
206    let value = values
207        .iter()
208        .next()?
209        .to_str()
210        .ok()?
211        .trim_matches([' ', '\t']);
212    if !value.bytes().all(|byte| byte.is_ascii_digit()) {
213        return None;
214    }
215    let seconds = value.parse::<u32>().ok()?;
216    (seconds > 0).then(|| Duration::from_secs(u64::from(seconds)))
217}
218
219/// Rejects frames from incompatible clients before interpreting their message.
220/// # Errors
221///
222/// Returns an error if the supplied value is invalid.
223pub fn validate_version(version: u16) -> Result<()> {
224    if version != PROTOCOL_VERSION {
225        return Err(Error::Protocol(format!(
226            "unsupported protocol version {version}; expected {PROTOCOL_VERSION}"
227        )));
228    }
229    Ok(())
230}
231
232pub(crate) fn validate_session_id(session_id: &str) -> Result<()> {
233    if session_id.trim().is_empty() || session_id.len() > 4 * 1024 {
234        return Err(Error::Config("session ID must be 1–4096 bytes".into()));
235    }
236    Ok(())
237}
238
239#[cfg(test)]
240mod tests {
241    use super::*;
242
243    #[tokio::test]
244    async fn completed_bulk_frames_do_not_pin_peak_memory_on_idle_connections() {
245        let payload = "x".repeat(MAX_FRAME_BYTES - 2);
246        let mut bytes = Vec::new();
247        write_frame(&mut bytes, &payload).await.unwrap();
248        write_frame(&mut bytes, &"next").await.unwrap();
249        let mut reader = FrameReader::new(bytes.as_slice());
250        assert_eq!(
251            read_frame::<String>(&mut reader).await.unwrap(),
252            Some(payload)
253        );
254        eprintln!(
255            "frame-reader retained bytes after 50 MiB frame: {}",
256            reader.buffer.capacity()
257        );
258        assert!(reader.buffer.capacity() <= 64 * 1024);
259        assert_eq!(
260            read_frame::<String>(&mut reader).await.unwrap().as_deref(),
261            Some("next")
262        );
263    }
264}