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;
14pub const MAX_HELLO_FRAME_LEN: usize = 0;
16pub const DEFAULT_MAX_FRAME_LEN: usize = 1024 * 1024;
18
19#[derive(Debug, Clone)]
24pub struct FrameCodec {
25 max_frame_len: usize,
26}
27
28impl FrameCodec {
29 #[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}