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