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 WorkerActivityDescriptor, 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 activity_descriptors_field_len(®istration.activities)?,
253 ])
254}
255
256fn activity_descriptors_field_len(
257 activities: &[WorkerActivityDescriptor],
258) -> Result<usize, ProtocolError> {
259 checked_u32_len(activities.len())?;
260 let mut total = U32_LEN;
261 for activity in activities {
262 total = sum_lengths(&[
263 total,
264 string_field_len(&activity.name)?,
265 string_field_len(&activity.input_schema_json)?,
266 string_field_len(&activity.output_schema_json)?,
267 ])?;
268 }
269 Ok(total)
270}
271
272fn worker_register_ack_payload_len(
273 outcome: &WorkerRegisterOutcome,
274) -> Result<usize, ProtocolError> {
275 match outcome {
276 WorkerRegisterOutcome::Accepted => Ok(U8_FIELD_LEN),
277 WorkerRegisterOutcome::Rejected { reason } => {
278 sum_lengths(&[U8_FIELD_LEN, string_field_len(reason)?])
279 }
280 }
281}
282
283fn write_handshake_payload(
284 frame: &Frame,
285 writer: &mut PayloadWriter<'_>,
286) -> Result<(), ProtocolError> {
287 match frame {
288 Frame::Connect {
289 min_version,
290 max_version,
291 auth_token,
292 ..
293 } => {
294 writer.write_slice(&min_version.to_wire_bytes())?;
295 writer.write_slice(&max_version.to_wire_bytes())?;
296 writer.write_bytes_field(auth_token)
297 }
298 Frame::ConnectAck {
299 selected_version,
300 capabilities,
301 ..
302 } => {
303 writer.write_slice(&selected_version.to_wire_bytes())?;
304 writer.write_u32(*capabilities)
305 }
306 _ => Err(ProtocolError::codec("frame type was not a handshake frame")),
307 }
308}
309
310fn write_pressure_payload(
311 frame: &Frame,
312 writer: &mut PayloadWriter<'_>,
313) -> Result<(), ProtocolError> {
314 match frame {
315 Frame::Accept {
316 referenced_message_id,
317 ..
318 } => writer.write_string_field(referenced_message_id.as_str()),
319 Frame::Defer {
320 referenced_message_id,
321 reason,
322 ..
323 }
324 | Frame::Reject {
325 referenced_message_id,
326 reason,
327 ..
328 } => {
329 writer.write_string_field(referenced_message_id.as_str())?;
330 writer.write_optional_string(reason.as_deref())
331 }
332 _ => Err(ProtocolError::codec("frame type was not a pressure frame")),
333 }
334}
335
336fn write_publish_payload(
337 frame: &Frame,
338 writer: &mut PayloadWriter<'_>,
339) -> Result<(), ProtocolError> {
340 match frame {
341 Frame::Publish {
342 channel,
343 envelope,
344 idempotency_key,
345 ..
346 } => {
347 writer.write_string_field(channel)?;
348 writer.write_bytes_field(&envelope.serialize()?)?;
349 if let Some(key) = idempotency_key {
354 writer.write_string_field(key)?;
355 }
356 Ok(())
357 }
358 _ => Err(ProtocolError::codec("frame type was not a publish frame")),
359 }
360}
361
362fn write_u64_prefixed_envelope(
365 writer: &mut PayloadWriter<'_>,
366 value: u64,
367 envelope: &MessageEnvelope,
368) -> Result<(), ProtocolError> {
369 writer.write_u64(value)?;
370 writer.write_bytes_field(&envelope.serialize()?)
371}
372
373fn write_push_payload(frame: &Frame, writer: &mut PayloadWriter<'_>) -> Result<(), ProtocolError> {
374 match frame {
375 Frame::Push {
376 correlation_id,
377 payload,
378 ..
379 }
380 | Frame::PushReply {
381 correlation_id,
382 payload,
383 ..
384 } => {
385 writer.write_u64(*correlation_id)?;
386 writer.write_bytes_field(payload)
387 }
388 _ => Err(ProtocolError::codec("frame type was not a push frame")),
389 }
390}
391
392fn write_worker_register_payload(
393 registration: &WorkerRegistration,
394 writer: &mut PayloadWriter<'_>,
395) -> Result<(), ProtocolError> {
396 writer.write_string_vec_field(®istration.namespaces)?;
397 writer.write_string_field(®istration.task_queue)?;
398 writer.write_optional_string(registration.node.as_deref())?;
401 writer.write_string_vec_field(®istration.activity_types)?;
402 writer.write_string_field(®istration.identity)?;
403 let count = u32::try_from(registration.activities.len())
404 .map_err(|_| ProtocolError::codec("activity descriptor count exceeded u32::MAX"))?;
405 writer.write_u32(count)?;
406 for activity in ®istration.activities {
407 writer.write_string_field(&activity.name)?;
408 writer.write_string_field(&activity.input_schema_json)?;
409 writer.write_string_field(&activity.output_schema_json)?;
410 }
411 Ok(())
412}
413
414fn write_worker_register_ack_payload(
415 outcome: &WorkerRegisterOutcome,
416 writer: &mut PayloadWriter<'_>,
417) -> Result<(), ProtocolError> {
418 match outcome {
419 WorkerRegisterOutcome::Accepted => writer.write_u8(WORKER_REGISTER_ACK_ACCEPTED),
420 WorkerRegisterOutcome::Rejected { reason } => {
421 writer.write_u8(WORKER_REGISTER_ACK_REJECTED)?;
422 writer.write_string_field(reason)
423 }
424 }
425}
426
427fn write_payload(frame: &Frame, buffer: &mut [u8]) -> Result<(), ProtocolError> {
428 let mut writer = PayloadWriter::new(buffer);
429 match frame {
430 Frame::Connect { .. } | Frame::ConnectAck { .. } => {
431 write_handshake_payload(frame, &mut writer)?;
432 }
433 Frame::ConnectError {
434 reason_code,
435 message,
436 ..
437 }
438 | Frame::SubscribeError {
439 reason_code,
440 message,
441 ..
442 }
443 | Frame::PublishError {
444 reason_code,
445 message,
446 ..
447 } => {
448 writer.write_u16(*reason_code)?;
449 writer.write_optional_string(message.as_deref())?;
450 }
451 Frame::Disconnect { .. } | Frame::Ping { .. } | Frame::Pong { .. } => {}
452 Frame::Subscribe {
453 channel,
454 accepted_schemas,
455 max_in_flight,
456 ..
457 } => {
458 writer.write_string_field(channel)?;
459 writer.write_schema_ids_field(accepted_schemas)?;
460 writer.write_u32(*max_in_flight)?;
461 }
462 Frame::SubscribeAck {
463 subscription_id,
464 selected_schema,
465 ..
466 } => {
467 writer.write_u64(*subscription_id)?;
468 writer.write_schema_id(*selected_schema)?;
469 }
470 Frame::Unsubscribe {
471 subscription_id, ..
472 } => writer.write_u64(*subscription_id)?,
473 Frame::Publish { .. } => write_publish_payload(frame, &mut writer)?,
474 Frame::PublishAck { message_id, .. } => writer.write_u64(*message_id)?,
475 Frame::ConversationOpen {
476 conversation_id,
477 subject,
478 ..
479 } => {
480 writer.write_u64(*conversation_id)?;
481 writer.write_string_field(subject)?;
482 }
483 Frame::ConversationMessage {
484 conversation_id,
485 envelope,
486 ..
487 } => write_u64_prefixed_envelope(&mut writer, *conversation_id, envelope)?,
488 Frame::ConversationClose {
489 conversation_id,
490 reason_code,
491 message,
492 ..
493 } => {
494 writer.write_u64(*conversation_id)?;
495 writer.write_optional_u16(*reason_code)?;
496 writer.write_optional_string(message.as_deref())?;
497 }
498 Frame::ConversationError {
499 conversation_id,
500 reason_code,
501 message,
502 ..
503 } => {
504 writer.write_u64(*conversation_id)?;
505 writer.write_u16(*reason_code)?;
506 writer.write_optional_string(message.as_deref())?;
507 }
508 Frame::Accept { .. } | Frame::Defer { .. } | Frame::Reject { .. } => {
509 write_pressure_payload(frame, &mut writer)?;
510 }
511 Frame::Push { .. } | Frame::PushReply { .. } => {
512 write_push_payload(frame, &mut writer)?;
513 }
514 Frame::Deliver {
515 delivery_seq,
516 envelope,
517 ..
518 } => write_u64_prefixed_envelope(&mut writer, *delivery_seq, envelope)?,
519 Frame::WorkerRegister { registration, .. } => {
520 write_worker_register_payload(registration, &mut writer)?;
521 }
522 Frame::WorkerRegisterAck { outcome, .. } => {
523 write_worker_register_ack_payload(outcome, &mut writer)?;
524 }
525 Frame::Unknown { payload, .. } => writer.write_slice(payload)?,
526 }
527 writer.finish()
528}
529
530fn decode_payload(
531 frame_type: FrameType,
532 flags: u8,
533 stream_id: u32,
534 payload: &[u8],
535) -> Result<Frame, ProtocolError> {
536 if let FrameType::Unknown(type_id) = frame_type {
537 return Ok(Frame::Unknown {
538 type_id,
539 flags,
540 stream_id,
541 payload: payload.to_vec(),
542 });
543 }
544
545 validate_stream(frame_type, stream_id)?;
546 decode_known_payload(frame_type, flags, stream_id, payload)
547}
548
549#[cfg(test)]
550mod tests;