mobius_gateway/wire/
codec.rs1use super::*;
2
3#[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
54pub struct FrameReader<R> {
56 reader: R,
57 buffer: Vec<u8>,
58}
59
60impl<R> FrameReader<R> {
61 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
86pub 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 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
143pub 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
219pub 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}