1mod known;
2mod payload;
3
4#[cfg(test)]
5mod tests_support;
6
7use super::causal::MessageId;
8use super::envelope::{MessageEnvelope, SchemaId};
9use super::error::ProtocolError;
10use super::frame::{
11 Frame, FrameType, HEADER_LEN, WORKER_REGISTER_ACK_ACCEPTED, WORKER_REGISTER_ACK_REJECTED,
12 WorkerRegisterOutcome, WorkerRegistration, validate_stream,
13};
14use super::version::ProtocolVersion;
15use known::decode_known_payload;
16use payload::{
17 PayloadReader, PayloadWriter, U16_LEN, U32_LEN, U64_LEN, bytes_field_len, checked_u32_len,
18 option_string_len, option_u16_len, schema_ids_field_len, string_field_len,
19 string_vec_field_len, sum_lengths,
20};
21
22const U8_FIELD_LEN: usize = 1;
26
27pub fn encoded_len(frame: &Frame) -> Result<usize, ProtocolError> {
34 frame.validate()?;
35 let payload_len = encoded_payload_len(frame)?;
36 HEADER_LEN
37 .checked_add(payload_len)
38 .ok_or_else(|| ProtocolError::codec("encoded frame length overflowed usize"))
39}
40
41pub fn encode(frame: &Frame, buffer: &mut [u8]) -> Result<usize, ProtocolError> {
53 frame.validate()?;
54 let payload_len = encoded_payload_len(frame)?;
55 let payload_length = u32::try_from(payload_len)
56 .map_err(|_| ProtocolError::codec("payload length exceeded u32::MAX"))?;
57 let total_len = HEADER_LEN
58 .checked_add(payload_len)
59 .ok_or_else(|| ProtocolError::codec("encoded frame length overflowed usize"))?;
60
61 if buffer.len() < total_len {
62 return Err(ProtocolError::codec("output buffer is too small"));
63 }
64
65 let Some(header) = buffer.get_mut(..HEADER_LEN) else {
66 return Err(ProtocolError::codec(
67 "output buffer is too small for header",
68 ));
69 };
70 write_header(frame, payload_length, header)?;
71
72 let Some(payload) = buffer.get_mut(HEADER_LEN..total_len) else {
73 return Err(ProtocolError::codec(
74 "output buffer is too small for payload",
75 ));
76 };
77 write_payload(frame, payload)?;
78
79 Ok(total_len)
80}
81
82pub fn decode(buffer: &[u8]) -> Result<(Frame, usize), ProtocolError> {
95 if buffer.len() < HEADER_LEN {
96 return Err(ProtocolError::IncompleteHeader {
97 message: Some("buffer shorter than fixed frame header".to_owned()),
98 });
99 }
100
101 let Some(header) = buffer.get(..HEADER_LEN) else {
102 return Err(ProtocolError::IncompleteHeader {
103 message: Some("buffer shorter than fixed frame header".to_owned()),
104 });
105 };
106 let mut header_reader = PayloadReader::new(header);
107 let type_id = header_reader.read_u8()?;
108 let flags = header_reader.read_u8()?;
109 let stream_id = header_reader.read_u32()?;
110 let payload_length = header_reader.read_u32()?;
111 header_reader.finish()?;
112
113 let payload_len = usize::try_from(payload_length)
114 .map_err(|_| ProtocolError::codec("payload length cannot fit usize"))?;
115 let total_len = HEADER_LEN
116 .checked_add(payload_len)
117 .ok_or_else(|| ProtocolError::codec("decoded frame length overflowed usize"))?;
118
119 if buffer.len() < total_len {
120 return Err(ProtocolError::TruncatedPayload {
121 message: Some("buffer shorter than declared payload length".to_owned()),
122 });
123 }
124
125 let Some(payload) = buffer.get(HEADER_LEN..total_len) else {
126 return Err(ProtocolError::TruncatedPayload {
127 message: Some("buffer shorter than declared payload length".to_owned()),
128 });
129 };
130
131 let frame_type = FrameType::from(type_id);
132 let frame = decode_payload(frame_type, flags, stream_id, payload)?;
133 Ok((frame, total_len))
134}
135
136fn write_header(
137 frame: &Frame,
138 payload_length: u32,
139 buffer: &mut [u8],
140) -> Result<(), ProtocolError> {
141 let mut writer = PayloadWriter::new(buffer);
142 writer.write_u8(u8::from(frame.frame_type()))?;
143 writer.write_u8(frame.flags())?;
144 writer.write_u32(frame.stream_id())?;
145 writer.write_u32(payload_length)?;
146 writer.finish()
147}
148
149fn encoded_payload_len(frame: &Frame) -> Result<usize, ProtocolError> {
150 match frame {
151 Frame::Connect { auth_token, .. } => sum_lengths(&[
152 ProtocolVersion::WIRE_LEN,
153 ProtocolVersion::WIRE_LEN,
154 bytes_field_len(auth_token)?,
155 ]),
156 Frame::ConnectAck { .. } => sum_lengths(&[ProtocolVersion::WIRE_LEN, U32_LEN]),
157 Frame::ConnectError { message, .. }
158 | Frame::SubscribeError { message, .. }
159 | Frame::PublishError { message, .. } => {
160 sum_lengths(&[U16_LEN, option_string_len(message.as_deref())?])
161 }
162 Frame::Disconnect { .. } | Frame::Ping { .. } | Frame::Pong { .. } => Ok(0),
163 Frame::Subscribe {
164 channel,
165 accepted_schemas,
166 ..
167 } => sum_lengths(&[
168 string_field_len(channel)?,
169 schema_ids_field_len(accepted_schemas)?,
170 U32_LEN,
171 ]),
172 Frame::SubscribeAck { .. } => sum_lengths(&[U64_LEN, SchemaId::WIRE_LEN]),
173 Frame::Unsubscribe { .. } | Frame::PublishAck { .. } => Ok(U64_LEN),
174 Frame::Publish {
175 channel,
176 envelope,
177 idempotency_key,
178 ..
179 } => {
180 let mut parts = vec![
181 string_field_len(channel)?,
182 envelope_bytes_field_len(envelope.encoded_len()?)?,
183 ];
184 if let Some(key) = idempotency_key {
185 parts.push(string_field_len(key)?);
186 }
187 sum_lengths(&parts)
188 }
189 Frame::ConversationOpen { subject, .. } => {
190 sum_lengths(&[U64_LEN, string_field_len(subject)?])
191 }
192 Frame::ConversationMessage { envelope, .. } | Frame::Deliver { envelope, .. } => {
196 sum_lengths(&[U64_LEN, envelope_bytes_field_len(envelope.encoded_len()?)?])
197 }
198 Frame::ConversationClose {
199 reason_code,
200 message,
201 ..
202 } => sum_lengths(&[
203 U64_LEN,
204 option_u16_len(*reason_code),
205 option_string_len(message.as_deref())?,
206 ]),
207 Frame::ConversationError { message, .. } => {
208 sum_lengths(&[U64_LEN, U16_LEN, option_string_len(message.as_deref())?])
209 }
210 Frame::Accept {
211 referenced_message_id,
212 ..
213 } => message_id_field_len(referenced_message_id),
214 Frame::Defer {
215 referenced_message_id,
216 reason,
217 ..
218 }
219 | Frame::Reject {
220 referenced_message_id,
221 reason,
222 ..
223 } => sum_lengths(&[
224 message_id_field_len(referenced_message_id)?,
225 option_string_len(reason.as_deref())?,
226 ]),
227 Frame::Push { payload, .. } | Frame::PushReply { payload, .. } => {
228 sum_lengths(&[U64_LEN, bytes_field_len(payload)?])
229 }
230 Frame::WorkerRegister { registration, .. } => worker_register_payload_len(registration),
231 Frame::WorkerRegisterAck { outcome, .. } => worker_register_ack_payload_len(outcome),
232 Frame::Unknown { payload, .. } => checked_u32_len(payload.len()).map(|()| payload.len()),
233 }
234}
235
236fn envelope_bytes_field_len(envelope_len: usize) -> Result<usize, ProtocolError> {
237 checked_u32_len(envelope_len)?;
238 sum_lengths(&[U32_LEN, envelope_len])
239}
240
241fn message_id_field_len(message_id: &MessageId) -> Result<usize, ProtocolError> {
242 string_field_len(message_id.as_str())
243}
244
245fn worker_register_payload_len(registration: &WorkerRegistration) -> Result<usize, ProtocolError> {
246 sum_lengths(&[
247 string_vec_field_len(®istration.namespaces)?,
248 string_field_len(®istration.task_queue)?,
249 option_string_len(registration.node.as_deref())?,
250 string_vec_field_len(®istration.activity_types)?,
251 string_field_len(®istration.identity)?,
252 ])
253}
254
255fn worker_register_ack_payload_len(
256 outcome: &WorkerRegisterOutcome,
257) -> Result<usize, ProtocolError> {
258 match outcome {
259 WorkerRegisterOutcome::Accepted => Ok(U8_FIELD_LEN),
260 WorkerRegisterOutcome::Rejected { reason } => {
261 sum_lengths(&[U8_FIELD_LEN, string_field_len(reason)?])
262 }
263 }
264}
265
266fn write_handshake_payload(
267 frame: &Frame,
268 writer: &mut PayloadWriter<'_>,
269) -> Result<(), ProtocolError> {
270 match frame {
271 Frame::Connect {
272 min_version,
273 max_version,
274 auth_token,
275 ..
276 } => {
277 writer.write_slice(&min_version.to_wire_bytes())?;
278 writer.write_slice(&max_version.to_wire_bytes())?;
279 writer.write_bytes_field(auth_token)
280 }
281 Frame::ConnectAck {
282 selected_version,
283 capabilities,
284 ..
285 } => {
286 writer.write_slice(&selected_version.to_wire_bytes())?;
287 writer.write_u32(*capabilities)
288 }
289 _ => Err(ProtocolError::codec("frame type was not a handshake frame")),
290 }
291}
292
293fn write_pressure_payload(
294 frame: &Frame,
295 writer: &mut PayloadWriter<'_>,
296) -> Result<(), ProtocolError> {
297 match frame {
298 Frame::Accept {
299 referenced_message_id,
300 ..
301 } => writer.write_string_field(referenced_message_id.as_str()),
302 Frame::Defer {
303 referenced_message_id,
304 reason,
305 ..
306 }
307 | Frame::Reject {
308 referenced_message_id,
309 reason,
310 ..
311 } => {
312 writer.write_string_field(referenced_message_id.as_str())?;
313 writer.write_optional_string(reason.as_deref())
314 }
315 _ => Err(ProtocolError::codec("frame type was not a pressure frame")),
316 }
317}
318
319fn write_publish_payload(
320 frame: &Frame,
321 writer: &mut PayloadWriter<'_>,
322) -> Result<(), ProtocolError> {
323 match frame {
324 Frame::Publish {
325 channel,
326 envelope,
327 idempotency_key,
328 ..
329 } => {
330 writer.write_string_field(channel)?;
331 writer.write_bytes_field(&envelope.serialize()?)?;
332 if let Some(key) = idempotency_key {
337 writer.write_string_field(key)?;
338 }
339 Ok(())
340 }
341 _ => Err(ProtocolError::codec("frame type was not a publish frame")),
342 }
343}
344
345fn write_u64_prefixed_envelope(
348 writer: &mut PayloadWriter<'_>,
349 value: u64,
350 envelope: &MessageEnvelope,
351) -> Result<(), ProtocolError> {
352 writer.write_u64(value)?;
353 writer.write_bytes_field(&envelope.serialize()?)
354}
355
356fn write_push_payload(frame: &Frame, writer: &mut PayloadWriter<'_>) -> Result<(), ProtocolError> {
357 match frame {
358 Frame::Push {
359 correlation_id,
360 payload,
361 ..
362 }
363 | Frame::PushReply {
364 correlation_id,
365 payload,
366 ..
367 } => {
368 writer.write_u64(*correlation_id)?;
369 writer.write_bytes_field(payload)
370 }
371 _ => Err(ProtocolError::codec("frame type was not a push frame")),
372 }
373}
374
375fn write_worker_register_payload(
376 registration: &WorkerRegistration,
377 writer: &mut PayloadWriter<'_>,
378) -> Result<(), ProtocolError> {
379 writer.write_string_vec_field(®istration.namespaces)?;
380 writer.write_string_field(®istration.task_queue)?;
381 writer.write_optional_string(registration.node.as_deref())?;
384 writer.write_string_vec_field(®istration.activity_types)?;
385 writer.write_string_field(®istration.identity)
386}
387
388fn write_worker_register_ack_payload(
389 outcome: &WorkerRegisterOutcome,
390 writer: &mut PayloadWriter<'_>,
391) -> Result<(), ProtocolError> {
392 match outcome {
393 WorkerRegisterOutcome::Accepted => writer.write_u8(WORKER_REGISTER_ACK_ACCEPTED),
394 WorkerRegisterOutcome::Rejected { reason } => {
395 writer.write_u8(WORKER_REGISTER_ACK_REJECTED)?;
396 writer.write_string_field(reason)
397 }
398 }
399}
400
401fn write_payload(frame: &Frame, buffer: &mut [u8]) -> Result<(), ProtocolError> {
402 let mut writer = PayloadWriter::new(buffer);
403 match frame {
404 Frame::Connect { .. } | Frame::ConnectAck { .. } => {
405 write_handshake_payload(frame, &mut writer)?;
406 }
407 Frame::ConnectError {
408 reason_code,
409 message,
410 ..
411 }
412 | Frame::SubscribeError {
413 reason_code,
414 message,
415 ..
416 }
417 | Frame::PublishError {
418 reason_code,
419 message,
420 ..
421 } => {
422 writer.write_u16(*reason_code)?;
423 writer.write_optional_string(message.as_deref())?;
424 }
425 Frame::Disconnect { .. } | Frame::Ping { .. } | Frame::Pong { .. } => {}
426 Frame::Subscribe {
427 channel,
428 accepted_schemas,
429 max_in_flight,
430 ..
431 } => {
432 writer.write_string_field(channel)?;
433 writer.write_schema_ids_field(accepted_schemas)?;
434 writer.write_u32(*max_in_flight)?;
435 }
436 Frame::SubscribeAck {
437 subscription_id,
438 selected_schema,
439 ..
440 } => {
441 writer.write_u64(*subscription_id)?;
442 writer.write_schema_id(*selected_schema)?;
443 }
444 Frame::Unsubscribe {
445 subscription_id, ..
446 } => writer.write_u64(*subscription_id)?,
447 Frame::Publish { .. } => write_publish_payload(frame, &mut writer)?,
448 Frame::PublishAck { message_id, .. } => writer.write_u64(*message_id)?,
449 Frame::ConversationOpen {
450 conversation_id,
451 subject,
452 ..
453 } => {
454 writer.write_u64(*conversation_id)?;
455 writer.write_string_field(subject)?;
456 }
457 Frame::ConversationMessage {
458 conversation_id,
459 envelope,
460 ..
461 } => write_u64_prefixed_envelope(&mut writer, *conversation_id, envelope)?,
462 Frame::ConversationClose {
463 conversation_id,
464 reason_code,
465 message,
466 ..
467 } => {
468 writer.write_u64(*conversation_id)?;
469 writer.write_optional_u16(*reason_code)?;
470 writer.write_optional_string(message.as_deref())?;
471 }
472 Frame::ConversationError {
473 conversation_id,
474 reason_code,
475 message,
476 ..
477 } => {
478 writer.write_u64(*conversation_id)?;
479 writer.write_u16(*reason_code)?;
480 writer.write_optional_string(message.as_deref())?;
481 }
482 Frame::Accept { .. } | Frame::Defer { .. } | Frame::Reject { .. } => {
483 write_pressure_payload(frame, &mut writer)?;
484 }
485 Frame::Push { .. } | Frame::PushReply { .. } => {
486 write_push_payload(frame, &mut writer)?;
487 }
488 Frame::Deliver {
489 delivery_seq,
490 envelope,
491 ..
492 } => write_u64_prefixed_envelope(&mut writer, *delivery_seq, envelope)?,
493 Frame::WorkerRegister { registration, .. } => {
494 write_worker_register_payload(registration, &mut writer)?;
495 }
496 Frame::WorkerRegisterAck { outcome, .. } => {
497 write_worker_register_ack_payload(outcome, &mut writer)?;
498 }
499 Frame::Unknown { payload, .. } => writer.write_slice(payload)?,
500 }
501 writer.finish()
502}
503
504fn decode_payload(
505 frame_type: FrameType,
506 flags: u8,
507 stream_id: u32,
508 payload: &[u8],
509) -> Result<Frame, ProtocolError> {
510 if let FrameType::Unknown(type_id) = frame_type {
511 return Ok(Frame::Unknown {
512 type_id,
513 flags,
514 stream_id,
515 payload: payload.to_vec(),
516 });
517 }
518
519 validate_stream(frame_type, stream_id)?;
520 decode_known_payload(frame_type, flags, stream_id, payload)
521}
522
523#[cfg(test)]
524mod tests;