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;
17
18#[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}