Skip to main content

relaygate_protocol/
codec.rs

1use bytes::{Buf, BufMut, Bytes, BytesMut};
2use tokio_util::codec::{Decoder, Encoder};
3use uuid::Uuid;
4
5use crate::{
6    BearerToken, BindingId, Destination, DestinationName, ErrorCode, Frame, Namespace,
7    PeerObservation, PipeId, ProtocolError, SessionId,
8};
9
10const MAGIC: [u8; 2] = *b"RG";
11const VERSION: u8 = 3;
12const HEADER_LEN: usize = 8;
13const MAX_STRING_LEN: usize = u16::MAX as usize;
14/// HELLO carries no credential or payload; PUBLISH and DIAL carry authorization.
15pub const MAX_HELLO_FRAME_LEN: usize = 0;
16/// Default payload limit a [`FrameCodec`] accepts, in bytes.
17pub const DEFAULT_MAX_FRAME_LEN: usize = 1024 * 1024;
18
19/// Length-delimited codec for SDK–Gateway [`Frame`] values.
20///
21/// The configured limit applies to the frame payload and excludes the fixed
22/// wire header.
23#[derive(Debug, Clone)]
24pub struct FrameCodec {
25    max_frame_len: usize,
26}
27
28impl FrameCodec {
29    /// Creates a codec that rejects payloads longer than `max_frame_len` bytes.
30    #[must_use]
31    pub const fn new(max_frame_len: usize) -> Self {
32        Self { max_frame_len }
33    }
34}
35
36impl Default for FrameCodec {
37    fn default() -> Self {
38        Self::new(DEFAULT_MAX_FRAME_LEN)
39    }
40}
41
42impl Encoder<Frame> for FrameCodec {
43    type Error = ProtocolError;
44
45    fn encode(&mut self, item: Frame, destination: &mut BytesMut) -> Result<(), Self::Error> {
46        let mut payload = BytesMut::new();
47        let kind = encode_payload(item, &mut payload)?;
48        if payload.len() > self.max_frame_len {
49            return Err(ProtocolError::FrameTooLarge {
50                actual: payload.len(),
51                maximum: self.max_frame_len,
52            });
53        }
54        let payload_len =
55            u32::try_from(payload.len()).map_err(|_| ProtocolError::LengthOverflow)?;
56        destination.reserve(HEADER_LEN + payload.len());
57        destination.extend_from_slice(&MAGIC);
58        destination.put_u8(VERSION);
59        destination.put_u8(kind);
60        destination.put_u32(payload_len);
61        destination.unsplit(payload);
62        Ok(())
63    }
64}
65
66impl Decoder for FrameCodec {
67    type Item = Frame;
68    type Error = ProtocolError;
69
70    fn decode(&mut self, source: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
71        if source.len() < HEADER_LEN {
72            return Ok(None);
73        }
74        if source[..2] != MAGIC {
75            return Err(ProtocolError::InvalidMagic);
76        }
77        let version = source[2];
78        if version != VERSION {
79            return Err(ProtocolError::UnsupportedVersion(version));
80        }
81        let kind = source[3];
82        let payload_len = u32::from_be_bytes([source[4], source[5], source[6], source[7]]) as usize;
83        if payload_len > self.max_frame_len {
84            return Err(ProtocolError::FrameTooLarge {
85                actual: payload_len,
86                maximum: self.max_frame_len,
87            });
88        }
89        if source.len() < HEADER_LEN + payload_len {
90            source.reserve(HEADER_LEN + payload_len - source.len());
91            return Ok(None);
92        }
93        source.advance(HEADER_LEN);
94        let payload = source.split_to(payload_len).freeze();
95        decode_payload(kind, payload).map(Some)
96    }
97}
98
99fn encode_payload(frame: Frame, destination: &mut BytesMut) -> Result<u8, ProtocolError> {
100    let kind = match frame {
101        Frame::Hello => 1,
102        Frame::Welcome { session_id } => {
103            put_session_id(destination, session_id);
104            2
105        }
106        Frame::SessionRejected { code, message } => {
107            destination.put_u8(code as u8);
108            put_string(destination, "message", &message)?;
109            3
110        }
111        Frame::Publish {
112            request_id,
113            destination: route_destination,
114            access_token,
115        } => {
116            destination.put_u64(request_id);
117            put_destination(destination, &route_destination)?;
118            put_string(destination, "access_token", access_token.expose_secret())?;
119            4
120        }
121        Frame::Published {
122            request_id,
123            binding_id,
124        } => {
125            destination.put_u64(request_id);
126            put_binding_id(destination, binding_id);
127            5
128        }
129        Frame::PublishFailed {
130            request_id,
131            code,
132            message,
133        } => {
134            destination.put_u64(request_id);
135            destination.put_u8(code as u8);
136            put_string(destination, "message", &message)?;
137            6
138        }
139        Frame::Unpublish {
140            request_id,
141            binding_id,
142        } => {
143            destination.put_u64(request_id);
144            put_binding_id(destination, binding_id);
145            7
146        }
147        Frame::Unpublished { request_id } => {
148            destination.put_u64(request_id);
149            8
150        }
151        Frame::Dial {
152            connection_id,
153            destination: route_destination,
154            access_token,
155        } => {
156            destination.put_u64(connection_id);
157            put_destination(destination, &route_destination)?;
158            put_string(destination, "access_token", access_token.expose_secret())?;
159            9
160        }
161        Frame::Offer {
162            pipe_id,
163            binding_id,
164            destination: route_destination,
165        } => {
166            put_pipe_id(destination, pipe_id);
167            put_binding_id(destination, binding_id);
168            put_destination(destination, &route_destination)?;
169            10
170        }
171        Frame::OfferAccepted { pipe_id } => {
172            put_pipe_id(destination, pipe_id);
173            11
174        }
175        Frame::OfferRejected {
176            pipe_id,
177            code,
178            message,
179        } => {
180            put_pipe_id(destination, pipe_id);
181            destination.put_u8(code as u8);
182            put_string(destination, "message", &message)?;
183            12
184        }
185        Frame::Opened { pipe_id } => {
186            put_pipe_id(destination, pipe_id);
187            13
188        }
189        Frame::DialFailed {
190            connection_id,
191            code,
192            observation,
193            message,
194        } => {
195            destination.put_u64(connection_id);
196            destination.put_u8(code as u8);
197            destination.put_u8(observation as u8);
198            put_string(destination, "message", &message)?;
199            14
200        }
201        Frame::Data { pipe_id, payload } => {
202            put_pipe_id(destination, pipe_id);
203            destination.extend_from_slice(&payload);
204            15
205        }
206        Frame::Fin { pipe_id } => {
207            put_pipe_id(destination, pipe_id);
208            16
209        }
210        Frame::Close { pipe_id } => {
211            put_pipe_id(destination, pipe_id);
212            17
213        }
214        Frame::Reset {
215            pipe_id,
216            code,
217            message,
218        } => {
219            put_pipe_id(destination, pipe_id);
220            destination.put_u8(code as u8);
221            put_string(destination, "message", &message)?;
222            18
223        }
224        Frame::Ping { nonce } => {
225            destination.put_u64(nonce);
226            19
227        }
228        Frame::Pong { nonce } => {
229            destination.put_u64(nonce);
230            20
231        }
232        Frame::Cancel { pipe_id } => {
233            put_pipe_id(destination, pipe_id);
234            21
235        }
236    };
237    Ok(kind)
238}
239
240fn decode_payload(kind: u8, payload: Bytes) -> Result<Frame, ProtocolError> {
241    let mut reader = PayloadReader::new(payload);
242    let frame = match kind {
243        1 => Frame::Hello,
244        2 => Frame::Welcome {
245            session_id: reader.session_id()?,
246        },
247        3 => Frame::SessionRejected {
248            code: reader.error_code()?,
249            message: reader.string("message")?,
250        },
251        4 => Frame::Publish {
252            request_id: reader.u64("request_id")?,
253            destination: reader.destination()?,
254            access_token: reader.access_token()?,
255        },
256        5 => Frame::Published {
257            request_id: reader.u64("request_id")?,
258            binding_id: reader.binding_id()?,
259        },
260        6 => Frame::PublishFailed {
261            request_id: reader.u64("request_id")?,
262            code: reader.error_code()?,
263            message: reader.string("message")?,
264        },
265        7 => Frame::Unpublish {
266            request_id: reader.u64("request_id")?,
267            binding_id: reader.binding_id()?,
268        },
269        8 => Frame::Unpublished {
270            request_id: reader.u64("request_id")?,
271        },
272        9 => Frame::Dial {
273            connection_id: reader.u64("connection_id")?,
274            destination: reader.destination()?,
275            access_token: reader.access_token()?,
276        },
277        10 => Frame::Offer {
278            pipe_id: reader.pipe_id()?,
279            binding_id: reader.binding_id()?,
280            destination: reader.destination()?,
281        },
282        11 => Frame::OfferAccepted {
283            pipe_id: reader.pipe_id()?,
284        },
285        12 => Frame::OfferRejected {
286            pipe_id: reader.pipe_id()?,
287            code: reader.error_code()?,
288            message: reader.string("message")?,
289        },
290        13 => Frame::Opened {
291            pipe_id: reader.pipe_id()?,
292        },
293        14 => Frame::DialFailed {
294            connection_id: reader.u64("connection_id")?,
295            code: reader.error_code()?,
296            observation: reader.observation()?,
297            message: reader.string("message")?,
298        },
299        15 => {
300            let pipe_id = reader.pipe_id()?;
301            let payload = reader.remaining();
302            Frame::Data { pipe_id, payload }
303        }
304        16 => Frame::Fin {
305            pipe_id: reader.pipe_id()?,
306        },
307        17 => Frame::Close {
308            pipe_id: reader.pipe_id()?,
309        },
310        18 => Frame::Reset {
311            pipe_id: reader.pipe_id()?,
312            code: reader.error_code()?,
313            message: reader.string("message")?,
314        },
315        19 => Frame::Ping {
316            nonce: reader.u64("nonce")?,
317        },
318        20 => Frame::Pong {
319            nonce: reader.u64("nonce")?,
320        },
321        21 => Frame::Cancel {
322            pipe_id: reader.pipe_id()?,
323        },
324        other => return Err(ProtocolError::UnknownFrameKind(other)),
325    };
326    reader.finish()?;
327    Ok(frame)
328}
329
330fn put_string(
331    destination: &mut BytesMut,
332    field: &'static str,
333    value: &str,
334) -> Result<(), ProtocolError> {
335    let length = value.len();
336    let wire_length = u16::try_from(length).map_err(|_| ProtocolError::FieldTooLong {
337        field,
338        actual: length,
339        maximum: MAX_STRING_LEN,
340    })?;
341    destination.put_u16(wire_length);
342    destination.extend_from_slice(value.as_bytes());
343    Ok(())
344}
345
346fn put_session_id(destination: &mut BytesMut, value: SessionId) {
347    destination.extend_from_slice(value.as_uuid().as_bytes());
348}
349
350fn put_binding_id(destination: &mut BytesMut, value: BindingId) {
351    destination.extend_from_slice(value.as_uuid().as_bytes());
352}
353
354fn put_destination(destination: &mut BytesMut, value: &Destination) -> Result<(), ProtocolError> {
355    put_string(destination, "namespace", value.namespace().as_str())?;
356    put_string(destination, "name", value.name().as_str())
357}
358
359fn put_pipe_id(destination: &mut BytesMut, value: PipeId) {
360    put_session_id(destination, value.origin_session_id());
361    destination.put_u64(value.connection_id());
362}
363
364struct PayloadReader {
365    payload: Bytes,
366    position: usize,
367}
368
369impl PayloadReader {
370    fn new(payload: Bytes) -> Self {
371        Self {
372            payload,
373            position: 0,
374        }
375    }
376
377    fn take(&mut self, length: usize, field: &'static str) -> Result<&[u8], ProtocolError> {
378        let end = self
379            .position
380            .checked_add(length)
381            .ok_or(ProtocolError::LengthOverflow)?;
382        let Some(bytes) = self.payload.get(self.position..end) else {
383            return Err(ProtocolError::Truncated(field));
384        };
385        self.position = end;
386        Ok(bytes)
387    }
388
389    fn u8(&mut self, field: &'static str) -> Result<u8, ProtocolError> {
390        Ok(self.take(1, field)?[0])
391    }
392
393    fn u16(&mut self, field: &'static str) -> Result<u16, ProtocolError> {
394        let bytes = self.take(2, field)?;
395        Ok(u16::from_be_bytes([bytes[0], bytes[1]]))
396    }
397
398    fn u64(&mut self, field: &'static str) -> Result<u64, ProtocolError> {
399        let bytes = self.take(8, field)?;
400        Ok(u64::from_be_bytes([
401            bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
402        ]))
403    }
404
405    fn string(&mut self, field: &'static str) -> Result<String, ProtocolError> {
406        let length = self.u16(field)? as usize;
407        let bytes = self.take(length, field)?;
408        let value = std::str::from_utf8(bytes).map_err(|_| ProtocolError::InvalidUtf8(field))?;
409        Ok(value.to_owned())
410    }
411
412    fn uuid(&mut self, field: &'static str) -> Result<Uuid, ProtocolError> {
413        let bytes = self.take(16, field)?;
414        Uuid::from_slice(bytes).map_err(|_| ProtocolError::Truncated(field))
415    }
416
417    fn session_id(&mut self) -> Result<SessionId, ProtocolError> {
418        self.uuid("session_id").map(SessionId::from_uuid)
419    }
420
421    fn binding_id(&mut self) -> Result<BindingId, ProtocolError> {
422        self.uuid("binding_id").map(BindingId::from_uuid)
423    }
424
425    fn destination(&mut self) -> Result<Destination, ProtocolError> {
426        let namespace = Namespace::new(&self.string("namespace")?)
427            .map_err(|_| ProtocolError::InvalidDestination)?;
428        let name = DestinationName::new(&self.string("name")?)
429            .map_err(|_| ProtocolError::InvalidDestination)?;
430        Ok(Destination::new(namespace, name))
431    }
432
433    fn access_token(&mut self) -> Result<BearerToken, ProtocolError> {
434        let length = self.u16("access_token")? as usize;
435        if length > crate::MAX_BEARER_TOKEN_BYTES {
436            return Err(ProtocolError::FieldTooLong {
437                field: "access_token",
438                actual: length,
439                maximum: crate::MAX_BEARER_TOKEN_BYTES,
440            });
441        }
442        let bytes = self.take(length, "access_token")?;
443        let value =
444            std::str::from_utf8(bytes).map_err(|_| ProtocolError::InvalidUtf8("access_token"))?;
445        BearerToken::new(value.to_owned())
446    }
447
448    fn pipe_id(&mut self) -> Result<PipeId, ProtocolError> {
449        let session_id = self.session_id()?;
450        let connection_id = self.u64("connection_id")?;
451        Ok(PipeId::new(session_id, connection_id))
452    }
453
454    fn error_code(&mut self) -> Result<ErrorCode, ProtocolError> {
455        let value = self.u8("error_code")?;
456        ErrorCode::from_wire(value).ok_or(ProtocolError::UnknownEnum {
457            name: "ErrorCode",
458            value,
459        })
460    }
461
462    fn observation(&mut self) -> Result<PeerObservation, ProtocolError> {
463        let value = self.u8("peer_observation")?;
464        PeerObservation::from_wire(value).ok_or(ProtocolError::UnknownEnum {
465            name: "PeerObservation",
466            value,
467        })
468    }
469
470    fn remaining(&mut self) -> Bytes {
471        let remaining = self.payload.slice(self.position..);
472        self.position = self.payload.len();
473        remaining
474    }
475
476    fn finish(self) -> Result<(), ProtocolError> {
477        let trailing = self.payload.len().saturating_sub(self.position);
478        if trailing == 0 {
479            Ok(())
480        } else {
481            Err(ProtocolError::TrailingBytes(trailing))
482        }
483    }
484}