1use std::convert::Infallible;
22
23use minicbor::data::Type;
24use minicbor::{Decoder, Encoder};
25use weida_core::Error;
26
27use crate::varint::{VarintError, decode_varint, encode_varint};
28
29pub mod limits {
33 pub const MAX_ENDPOINT_BYTES: usize = 512;
35 pub const MAX_CONTENT_TYPE_BYTES: usize = 256;
37 pub const MAX_TRACEPARENT_BYTES: usize = 128;
39 pub const MAX_TRACESTATE_BYTES: usize = 512;
41 pub const MAX_TOPIC_BYTES: usize = 256;
43 pub const PRODUCER_BYTES: usize = 32;
49 pub const MAX_FILTER_BYTES: usize = 256;
51 pub const MAX_MESSAGE_BYTES: usize = 1024;
53 pub const MAX_LIST_ITEMS: usize = 64;
58 pub const MAX_REPORT_LEVELS: usize = 16;
66 pub const MAX_SKIP_DEPTH: usize = 8;
68 pub const MAX_LAYER: u8 = 15;
71 pub const MAX_PATH_RECORD_BYTES: usize = 256;
73}
74
75mod hello_key {
77 pub const VERSIONS: u64 = 0;
78 pub const MAX_HEADER_BYTES: u64 = 1;
79 pub const MAX_TRANSFERS: u64 = 2;
80 pub const CAPABILITIES: u64 = 3;
81 pub const REQUIRED_CAPABILITIES: u64 = 4;
82 pub const GUARANTEES_OFFERED: u64 = 5;
83 pub const GUARANTEES_REQUIRED: u64 = 6;
84}
85
86mod data_key {
88 pub const ENDPOINT: u64 = 0;
89 pub const CONTENT_LEN: u64 = 1;
90 pub const CONTENT_TYPE: u64 = 2;
91 pub const TRACEPARENT: u64 = 3;
92 pub const TRACESTATE: u64 = 4;
93 pub const TOPIC: u64 = 5;
94 pub const SEQUENCE: u64 = 6;
95 pub const PRODUCER: u64 = 7;
96 pub const ACHIEVED: u64 = 8;
97 pub const REPORT_ID: u64 = 9;
98 pub const REPORT: u64 = 10;
99 pub const REPORT_MODE: u64 = 11;
100 pub const SEGMENT: u64 = 13;
103 pub const LAYER: u64 = 14;
104}
105
106mod error_key {
108 pub const CODE: u64 = 0;
109 pub const MESSAGE: u64 = 1;
110}
111
112mod subscription_key {
114 pub const ENDPOINT: u64 = 0;
115 pub const FILTER: u64 = 1;
116 pub const MAX_AGE_MS: u64 = 2;
117 pub const MAX_LAYER: u64 = 3;
118}
119
120mod credit_key {
122 pub const ENDPOINT: u64 = 0;
123 pub const FILTER: u64 = 1;
124 pub const LIMIT: u64 = 2;
125}
126
127mod cursor_key {
129 pub const REPORT_ID: u64 = 0;
130}
131
132mod flow_key {
134 pub const ENDPOINT: u64 = 0;
135 pub const FLOW: u64 = 1;
136 pub const CONTENT_TYPE: u64 = 2;
137 pub const TRACEPARENT: u64 = 3;
138 pub const TRACESTATE: u64 = 4;
139 pub const TOPIC: u64 = 5;
140}
141
142pub mod filter {
157 use super::HeaderError;
158
159 pub const SEPARATOR: char = '.';
161 pub const ONE_SEGMENT: &str = "*";
163 pub const REST: &str = "#";
165
166 pub fn matches(topic: &str, filter: &str) -> bool {
188 if filter.is_empty() {
189 return true;
190 }
191 let mut topic_segments = topic.split(SEPARATOR);
192 let mut filter_segments = filter.split(SEPARATOR);
193 loop {
194 let Some(pattern) = filter_segments.next() else {
195 return topic_segments.next().is_none();
197 };
198 if pattern == REST {
201 return true;
202 }
203 let Some(segment) = topic_segments.next() else {
204 return false;
205 };
206 if pattern != ONE_SEGMENT && pattern != segment {
207 return false;
208 }
209 }
210 }
211
212 pub fn validate(filter: &str) -> Result<(), HeaderError> {
217 let mut segments = filter.split(SEPARATOR).peekable();
218 while let Some(segment) = segments.next() {
219 let is_last = segments.peek().is_none();
220 if segment.contains(ONE_SEGMENT) && segment != ONE_SEGMENT {
221 return Err(HeaderError::InvalidFilter(
222 "`*` must occupy a whole segment",
223 ));
224 }
225 if segment.contains(REST) {
226 if segment != REST {
227 return Err(HeaderError::InvalidFilter(
228 "`#` must occupy a whole segment",
229 ));
230 }
231 if !is_last {
232 return Err(HeaderError::InvalidFilter("`#` must be the final segment"));
233 }
234 }
235 }
236 Ok(())
237 }
238}
239
240mod guarantee_key {
242 pub const DELIVERY: u64 = 0;
243 pub const ACKNOWLEDGEMENT: u64 = 1;
244 pub const DURABILITY: u64 = 2;
245 pub const REPLICAS: u64 = 3;
246 pub const ORDERING: u64 = 4;
247 pub const DEDUPLICATION: u64 = 5;
248 pub const DEDUP_WINDOW_MS: u64 = 6;
249 pub const BACKPRESSURE: u64 = 7;
250 pub const PRODUCER_NAMING: u64 = 8;
251 pub const CONTROL_ISOLATED: u64 = 9;
252}
253
254macro_rules! wire_enum {
265 ($(#[$meta:meta])* $name:ident { $($(#[$vmeta:meta])* $variant:ident = $value:literal),+ $(,)? }) => {
266 $(#[$meta])*
267 #[derive(Clone, Copy, Debug, Default, PartialEq, Eq, PartialOrd, Ord, Hash)]
268 pub enum $name {
269 $($(#[$vmeta])* $variant,)+
270 }
271
272 impl $name {
273 pub fn to_wire(self) -> u64 {
275 match self {
276 $($name::$variant => $value,)+
277 }
278 }
279
280 pub fn from_wire(value: u64) -> Option<$name> {
283 match value {
284 $($value => Some($name::$variant),)+
285 _ => None,
286 }
287 }
288 }
289
290 const _: () = {
295 let values = [$($value as u64),+];
296 let mut i = 1;
297 while i < values.len() {
298 assert!(
299 values[i - 1] < values[i],
300 concat!(
301 stringify!($name),
302 ": wire values must ascend with declaration order, ",
303 "because the derived Ord is the ladder"
304 )
305 );
306 i += 1;
307 }
308 };
309 };
310}
311
312wire_enum! {
313 Delivery {
317 #[default]
319 BestEffort = 0,
320 AtMostOnce = 1,
322 AtLeastOnce = 2,
324 }
325}
326
327wire_enum! {
328 Acknowledgement {
331 None = 0,
333 #[default]
335 TransportReceipt = 1,
336 Accepted = 2,
338 Stored = 3,
340 Replicated = 4,
342 Processed = 5,
344 }
345}
346
347wire_enum! {
348 Durability {
351 #[default]
353 Written = 0,
354 Flushed = 1,
356 }
357}
358
359wire_enum! {
360 OrderingMode {
362 #[default]
364 None = 0,
365 PerProducerDetect = 1,
367 PerProducerReassemble = 2,
369 PerKey = 3,
371 Total = 4,
373 }
374}
375
376wire_enum! {
377 Deduplication {
379 #[default]
381 None = 0,
382 Bounded = 1,
384 Durable = 2,
386 }
387}
388
389wire_enum! {
390 Backpressure {
393 #[default]
395 Block = 0,
396 Reject = 1,
398 Drop = 2,
400 Spill = 3,
402 Coalesce = 4,
404 }
405}
406
407wire_enum! {
408 ProducerNaming {
412 #[default]
416 Fingerprint = 0,
417 Stable = 1,
419 }
420}
421
422wire_enum! {
423 ReportMode {
432 #[default]
435 Progress = 0,
436 FinalOnly = 1,
438 }
439}
440
441#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
452pub enum CursorLevel {
453 Known(Acknowledgement),
455 Application(u64),
457}
458
459impl CursorLevel {
460 pub const APPLICATION_FLOOR: u64 = 16;
462
463 pub fn to_wire(self) -> u64 {
465 match self {
466 CursorLevel::Known(level) => level.to_wire(),
467 CursorLevel::Application(value) => value,
468 }
469 }
470
471 pub fn from_wire(value: u64) -> Option<CursorLevel> {
474 if value >= CursorLevel::APPLICATION_FLOOR {
475 Some(CursorLevel::Application(value))
476 } else {
477 Acknowledgement::from_wire(value).map(CursorLevel::Known)
478 }
479 }
480
481 pub fn application(value: u64) -> Option<CursorLevel> {
484 (value >= CursorLevel::APPLICATION_FLOOR).then_some(CursorLevel::Application(value))
485 }
486}
487
488#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
497pub struct GuaranteeSet {
498 pub delivery: Delivery,
500 pub acknowledgement: Acknowledgement,
502 pub durability: Option<Durability>,
504 pub replicas: Option<u64>,
506 pub ordering: OrderingMode,
508 pub deduplication: Deduplication,
510 pub dedup_window_ms: Option<u64>,
512 pub backpressure: Backpressure,
514 pub producer_naming: ProducerNaming,
516 pub control_isolated: bool,
520}
521
522impl GuaranteeSet {
523 pub const CORE: GuaranteeSet = GuaranteeSet {
525 delivery: Delivery::BestEffort,
526 acknowledgement: Acknowledgement::TransportReceipt,
527 durability: None,
528 replicas: None,
529 ordering: OrderingMode::None,
530 deduplication: Deduplication::None,
531 dedup_window_ms: None,
532 backpressure: Backpressure::Block,
533 producer_naming: ProducerNaming::Fingerprint,
534 control_isolated: false,
535 };
536
537 pub fn is_core(&self) -> bool {
540 *self == GuaranteeSet::CORE
541 }
542
543 fn validate(&self) -> Result<(), HeaderError> {
545 let stored_or_replicated = matches!(
546 self.acknowledgement,
547 Acknowledgement::Stored | Acknowledgement::Replicated
548 );
549 if self.durability.is_some() && !stored_or_replicated {
550 return Err(HeaderError::InvalidGuarantees(
551 "durability without Stored or Replicated",
552 ));
553 }
554 match self.replicas {
555 Some(_) if self.acknowledgement != Acknowledgement::Replicated => {
556 return Err(HeaderError::InvalidGuarantees(
557 "replicas without Replicated",
558 ));
559 }
560 Some(n) if n < 2 => {
561 return Err(HeaderError::InvalidGuarantees(
562 "a replica count below 2 is not a replication",
563 ));
564 }
565 _ => {}
566 }
567 match (self.deduplication, self.dedup_window_ms) {
568 (Deduplication::Bounded, None) => {
569 return Err(HeaderError::InvalidGuarantees(
570 "Bounded deduplication without a window",
571 ));
572 }
573 (level, Some(_)) if level != Deduplication::Bounded => {
574 return Err(HeaderError::InvalidGuarantees(
575 "a dedup window without Bounded deduplication",
576 ));
577 }
578 _ => {}
579 }
580 Ok(())
581 }
582
583 fn encode_into(
585 &self,
586 e: &mut Encoder<Vec<u8>>,
587 ) -> Result<(), minicbor::encode::Error<Infallible>> {
588 let core = GuaranteeSet::CORE;
589 let count = u64::from(self.delivery != core.delivery)
590 + u64::from(self.acknowledgement != core.acknowledgement)
591 + u64::from(self.durability.is_some())
592 + u64::from(self.replicas.is_some())
593 + u64::from(self.ordering != core.ordering)
594 + u64::from(self.deduplication != core.deduplication)
595 + u64::from(self.dedup_window_ms.is_some())
596 + u64::from(self.backpressure != core.backpressure)
597 + u64::from(self.producer_naming != core.producer_naming)
598 + u64::from(self.control_isolated != core.control_isolated);
599 e.map(count)?;
600 if self.delivery != core.delivery {
601 e.u64(guarantee_key::DELIVERY)?
602 .u64(self.delivery.to_wire())?;
603 }
604 if self.acknowledgement != core.acknowledgement {
605 e.u64(guarantee_key::ACKNOWLEDGEMENT)?
606 .u64(self.acknowledgement.to_wire())?;
607 }
608 if let Some(durability) = self.durability {
609 e.u64(guarantee_key::DURABILITY)?
610 .u64(durability.to_wire())?;
611 }
612 if let Some(replicas) = self.replicas {
613 e.u64(guarantee_key::REPLICAS)?.u64(replicas)?;
614 }
615 if self.ordering != core.ordering {
616 e.u64(guarantee_key::ORDERING)?
617 .u64(self.ordering.to_wire())?;
618 }
619 if self.deduplication != core.deduplication {
620 e.u64(guarantee_key::DEDUPLICATION)?
621 .u64(self.deduplication.to_wire())?;
622 }
623 if let Some(window) = self.dedup_window_ms {
624 e.u64(guarantee_key::DEDUP_WINDOW_MS)?.u64(window)?;
625 }
626 if self.backpressure != core.backpressure {
627 e.u64(guarantee_key::BACKPRESSURE)?
628 .u64(self.backpressure.to_wire())?;
629 }
630 if self.producer_naming != core.producer_naming {
631 e.u64(guarantee_key::PRODUCER_NAMING)?
632 .u64(self.producer_naming.to_wire())?;
633 }
634 if self.control_isolated != core.control_isolated {
635 e.u64(guarantee_key::CONTROL_ISOLATED)?
636 .u64(u64::from(self.control_isolated))?;
637 }
638 Ok(())
639 }
640
641 fn decode_from(m: &mut MapReader<'_, '_>) -> Result<GuaranteeSet, HeaderError> {
646 let mut set = GuaranteeSet::CORE;
647 let mut inner = MapReader::new(m.d)?;
648 while let Some(key) = inner.next_key()? {
649 match key {
650 guarantee_key::DELIVERY => set.delivery = level(inner.u64()?, "delivery")?,
651 guarantee_key::ACKNOWLEDGEMENT => {
652 set.acknowledgement = level(inner.u64()?, "acknowledgement")?;
653 }
654 guarantee_key::DURABILITY => {
655 set.durability = Some(level(inner.u64()?, "durability")?);
656 }
657 guarantee_key::REPLICAS => set.replicas = Some(inner.u64()?),
658 guarantee_key::ORDERING => set.ordering = level(inner.u64()?, "ordering")?,
659 guarantee_key::DEDUPLICATION => {
660 set.deduplication = level(inner.u64()?, "deduplication")?;
661 }
662 guarantee_key::DEDUP_WINDOW_MS => set.dedup_window_ms = Some(inner.u64()?),
663 guarantee_key::BACKPRESSURE => {
664 set.backpressure = level(inner.u64()?, "backpressure")?;
665 }
666 guarantee_key::PRODUCER_NAMING => {
667 set.producer_naming = level(inner.u64()?, "producer naming")?;
668 }
669 guarantee_key::CONTROL_ISOLATED => {
670 set.control_isolated = match inner.u64()? {
671 0 => false,
672 1 => true,
673 _ => {
674 return Err(HeaderError::InvalidGuarantees(
675 "control_isolated is 0 or 1",
676 ));
677 }
678 };
679 }
680 _ => inner.skip()?,
681 }
682 }
683 set.validate()?;
684 Ok(set)
685 }
686
687 pub fn intersect(&self, other: &GuaranteeSet) -> Result<GuaranteeSet, &'static str> {
696 if self.backpressure != other.backpressure {
697 return Err("backpressure");
698 }
699 if self.producer_naming != other.producer_naming {
700 return Err("producer naming");
701 }
702 if self.durability.is_some()
703 && other.durability.is_some()
704 && self.durability != other.durability
705 {
706 return Err("durability");
707 }
708 if self.replicas.is_some() && other.replicas.is_some() && self.replicas != other.replicas {
709 return Err("replicas");
710 }
711
712 let acknowledgement = self.acknowledgement.min(other.acknowledgement);
713 let keeps_durability = matches!(
714 acknowledgement,
715 Acknowledgement::Stored | Acknowledgement::Replicated
716 );
717 let deduplication = self.deduplication.min(other.deduplication);
718 let mut merged = GuaranteeSet {
719 delivery: self.delivery.min(other.delivery),
720 acknowledgement,
721 durability: keeps_durability
725 .then_some(self.durability.or(other.durability))
726 .flatten(),
727 replicas: (acknowledgement == Acknowledgement::Replicated)
728 .then_some(self.replicas.or(other.replicas))
729 .flatten(),
730 ordering: self.ordering.min(other.ordering),
731 deduplication,
732 dedup_window_ms: None,
734 backpressure: self.backpressure,
735 producer_naming: self.producer_naming,
736 control_isolated: self.control_isolated && other.control_isolated,
737 };
738 if deduplication == Deduplication::Bounded {
739 merged.dedup_window_ms = match (self.dedup_window_ms, other.dedup_window_ms) {
740 (Some(a), Some(b)) => Some(a.min(b)),
741 (Some(a), None) | (None, Some(a)) => Some(a),
742 (None, None) => None,
743 };
744 }
745 Ok(merged)
746 }
747
748 pub fn reaches(&self, required: &GuaranteeSet) -> bool {
754 if self.delivery < required.delivery
755 || self.acknowledgement < required.acknowledgement
756 || self.ordering < required.ordering
757 || self.deduplication < required.deduplication
758 {
759 return false;
760 }
761 if self.backpressure != required.backpressure
762 || self.producer_naming != required.producer_naming
763 {
764 return false;
765 }
766 if !self.control_isolated && required.control_isolated {
767 return false;
768 }
769 match (self.durability, required.durability) {
770 (_, None) => {}
771 (Some(have), Some(want)) if have >= want => {}
772 _ => return false,
773 }
774 match (self.replicas, required.replicas) {
775 (_, None) => {}
776 (Some(have), Some(want)) if have >= want => {}
777 _ => return false,
778 }
779 match (self.dedup_window_ms, required.dedup_window_ms) {
780 (_, None) => {}
781 (Some(have), Some(want)) if have >= want => {}
782 _ => return false,
783 }
784 true
785 }
786}
787
788fn level<T: WireLevel>(value: u64, dimension: &'static str) -> Result<T, HeaderError> {
790 T::from_wire_value(value).ok_or(HeaderError::UnknownLevel { dimension, value })
791}
792
793fn layer(value: u64, reason: &'static str) -> Result<u8, HeaderError> {
796 u8::try_from(value)
797 .ok()
798 .filter(|layer| *layer <= limits::MAX_LAYER)
799 .ok_or(HeaderError::InvalidLayer(reason))
800}
801
802trait WireLevel: Sized {
804 fn from_wire_value(value: u64) -> Option<Self>;
805}
806
807macro_rules! impl_wire_level {
808 ($($name:ident),+ $(,)?) => {
809 $(impl WireLevel for $name {
810 fn from_wire_value(value: u64) -> Option<$name> {
811 $name::from_wire(value)
812 }
813 })+
814 };
815}
816
817impl_wire_level!(
818 Delivery,
819 Acknowledgement,
820 Durability,
821 OrderingMode,
822 Deduplication,
823 Backpressure,
824 ProducerNaming,
825 ReportMode,
826);
827
828#[derive(Clone, Debug, PartialEq, Eq)]
830pub enum HeaderError {
831 Malformed(&'static str),
833 Indefinite,
835 DuplicateKey(u64),
837 UnorderedKey(u64),
839 NonUintKey,
841 MissingKey(u64),
843 StringTooLong {
845 key: u64,
847 len: usize,
849 max: usize,
851 },
852 ListTooLong {
854 key: u64,
856 len: u64,
858 max: usize,
860 },
861 DepthExceeded,
863 TrailingBytes,
865 UnknownLevel {
868 dimension: &'static str,
870 value: u64,
872 },
873 InvalidGuarantees(&'static str),
876 InvalidFilter(&'static str),
878 InvalidReport(&'static str),
882 InvalidLayer(&'static str),
886 InvalidPathReport(&'static str),
889}
890
891impl std::fmt::Display for HeaderError {
892 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
893 match self {
894 HeaderError::Malformed(what) => write!(f, "malformed header: {what}"),
895 HeaderError::Indefinite => f.write_str("indefinite-length items are not allowed"),
896 HeaderError::DuplicateKey(k) => write!(f, "duplicate header key {k}"),
897 HeaderError::UnorderedKey(k) => {
898 write!(f, "header key {k} is out of ascending order")
899 }
900 HeaderError::NonUintKey => f.write_str("header key is not an unsigned integer"),
901 HeaderError::MissingKey(k) => write!(f, "required header key {k} is missing"),
902 HeaderError::StringTooLong { key, len, max } => {
903 write!(
904 f,
905 "key {key}: text of {len} bytes exceeds the {max} byte cap"
906 )
907 }
908 HeaderError::ListTooLong { key, len, max } => {
909 write!(
910 f,
911 "key {key}: list of {len} items exceeds the {max} item cap"
912 )
913 }
914 HeaderError::DepthExceeded => f.write_str("unknown field nested too deeply"),
915 HeaderError::TrailingBytes => f.write_str("trailing bytes after the header"),
916 HeaderError::UnknownLevel { dimension, value } => {
917 write!(f, "unknown {dimension} level {value}")
918 }
919 HeaderError::InvalidGuarantees(why) => write!(f, "invalid guarantee set: {why}"),
920 HeaderError::InvalidFilter(why) => write!(f, "invalid topic filter: {why}"),
921 HeaderError::InvalidReport(reason) => write!(f, "invalid report: {reason}"),
922 HeaderError::InvalidLayer(reason) => write!(f, "invalid layer: {reason}"),
923 HeaderError::InvalidPathReport(reason) => write!(f, "invalid path report: {reason}"),
924 }
925 }
926}
927
928impl std::error::Error for HeaderError {}
929
930impl From<HeaderError> for Error {
931 fn from(e: HeaderError) -> Error {
932 Error::Protocol(e.to_string())
933 }
934}
935
936fn encode_into_with(
944 out: &mut Vec<u8>,
945 f: impl FnOnce(&mut Encoder<Vec<u8>>) -> Result<(), minicbor::encode::Error<Infallible>>,
946) {
947 let mut e = Encoder::new(std::mem::take(out));
951 f(&mut e).expect("encoding into a Vec is infallible");
952 *out = e.into_writer();
953}
954
955fn skip_value(d: &mut Decoder<'_>, max_depth: usize) -> Result<(), HeaderError> {
962 let mut stack: Vec<u64> = Vec::new();
965 let mut remaining: u64 = 1;
966
967 loop {
968 if remaining == 0 {
969 match stack.pop() {
970 Some(outer) => {
971 remaining = outer;
972 continue;
973 }
974 None => return Ok(()),
975 }
976 }
977 remaining -= 1;
978
979 let ty = d
980 .datatype()
981 .map_err(|_| HeaderError::Malformed("truncated value"))?;
982 let nested = match ty {
983 Type::Bool => {
984 d.bool().map_err(|_| HeaderError::Malformed("bool"))?;
985 None
986 }
987 Type::Null => {
988 d.null().map_err(|_| HeaderError::Malformed("null"))?;
989 None
990 }
991 Type::Undefined => {
992 d.undefined()
993 .map_err(|_| HeaderError::Malformed("undefined"))?;
994 None
995 }
996 Type::U8
997 | Type::U16
998 | Type::U32
999 | Type::U64
1000 | Type::I8
1001 | Type::I16
1002 | Type::I32
1003 | Type::I64
1004 | Type::Int => {
1005 d.int().map_err(|_| HeaderError::Malformed("integer"))?;
1006 None
1007 }
1008 Type::F32 | Type::F64 => {
1009 d.f64().map_err(|_| HeaderError::Malformed("float"))?;
1010 None
1011 }
1012 Type::Bytes => {
1013 d.bytes()
1014 .map_err(|_| HeaderError::Malformed("byte string"))?;
1015 None
1016 }
1017 Type::String => {
1018 d.str().map_err(|_| HeaderError::Malformed("text string"))?;
1019 None
1020 }
1021 Type::Array => Some(
1022 d.array()
1023 .map_err(|_| HeaderError::Malformed("array"))?
1024 .ok_or(HeaderError::Indefinite)?,
1025 ),
1026 Type::Map => {
1027 let pairs = d
1028 .map()
1029 .map_err(|_| HeaderError::Malformed("map"))?
1030 .ok_or(HeaderError::Indefinite)?;
1031 Some(
1032 pairs
1033 .checked_mul(2)
1034 .ok_or(HeaderError::Malformed("map length overflow"))?,
1035 )
1036 }
1037 Type::BytesIndef | Type::StringIndef | Type::ArrayIndef | Type::MapIndef => {
1038 return Err(HeaderError::Indefinite);
1039 }
1040 Type::Break => return Err(HeaderError::Malformed("unexpected break")),
1041 Type::Tag => return Err(HeaderError::Malformed("tags are not allowed")),
1044 Type::F16 => return Err(HeaderError::Malformed("half floats are not allowed")),
1045 Type::Simple => return Err(HeaderError::Malformed("simple values are not allowed")),
1046 Type::Unknown(_) => return Err(HeaderError::Malformed("unknown major type")),
1047 };
1048
1049 if let Some(count) = nested
1050 && count > 0
1051 {
1052 if stack.len() >= max_depth {
1053 return Err(HeaderError::DepthExceeded);
1054 }
1055 stack.push(remaining);
1056 remaining = count;
1057 }
1058 }
1059}
1060
1061struct MapReader<'a, 'b> {
1063 d: &'a mut Decoder<'b>,
1064 remaining: u64,
1065 seen: u64,
1068 last: Option<u64>,
1075}
1076
1077impl<'a, 'b> MapReader<'a, 'b> {
1078 fn new(d: &'a mut Decoder<'b>) -> Result<MapReader<'a, 'b>, HeaderError> {
1079 let len = d
1080 .map()
1081 .map_err(|_| HeaderError::Malformed("header is not a map"))?
1082 .ok_or(HeaderError::Indefinite)?;
1083 Ok(MapReader {
1084 d,
1085 remaining: len,
1086 seen: 0,
1087 last: None,
1088 })
1089 }
1090
1091 fn next_key(&mut self) -> Result<Option<u64>, HeaderError> {
1092 if self.remaining == 0 {
1093 return Ok(None);
1094 }
1095 self.remaining -= 1;
1096 match self.d.datatype() {
1097 Ok(Type::U8 | Type::U16 | Type::U32 | Type::U64) => {}
1098 Ok(_) => return Err(HeaderError::NonUintKey),
1099 Err(_) => return Err(HeaderError::Malformed("truncated key")),
1100 }
1101 let key = self.d.u64().map_err(|_| HeaderError::NonUintKey)?;
1102 if let Some(prev) = self.last {
1103 if key == prev {
1104 return Err(HeaderError::DuplicateKey(key));
1105 }
1106 if key < prev {
1107 return Err(HeaderError::UnorderedKey(key));
1108 }
1109 }
1110 self.last = Some(key);
1111 if key < 64 {
1112 self.seen |= 1u64 << key;
1113 }
1114 Ok(Some(key))
1115 }
1116
1117 fn saw(&self, key: u64) -> bool {
1118 key < 64 && self.seen & (1u64 << key) != 0
1119 }
1120
1121 fn require(&self, key: u64) -> Result<(), HeaderError> {
1122 if self.saw(key) {
1123 Ok(())
1124 } else {
1125 Err(HeaderError::MissingKey(key))
1126 }
1127 }
1128
1129 fn u64(&mut self) -> Result<u64, HeaderError> {
1130 self.d
1131 .u64()
1132 .map_err(|_| HeaderError::Malformed("expected an unsigned integer"))
1133 }
1134
1135 fn text(&mut self, key: u64, max: usize) -> Result<String, HeaderError> {
1136 let s = self
1137 .d
1138 .str()
1139 .map_err(|_| HeaderError::Malformed("expected a text string"))?;
1140 if s.len() > max {
1141 return Err(HeaderError::StringTooLong {
1142 key,
1143 len: s.len(),
1144 max,
1145 });
1146 }
1147 Ok(s.to_owned())
1148 }
1149
1150 fn byte_array<const N: usize>(&mut self, key: u64) -> Result<[u8; N], HeaderError> {
1157 let bytes = self
1158 .d
1159 .bytes()
1160 .map_err(|_| HeaderError::Malformed("expected a byte string"))?;
1161 if bytes.len() > N {
1162 return Err(HeaderError::StringTooLong {
1163 key,
1164 len: bytes.len(),
1165 max: N,
1166 });
1167 }
1168 bytes
1169 .try_into()
1170 .map_err(|_| HeaderError::Malformed("byte string has the wrong length"))
1171 }
1172
1173 fn uint_list(&mut self, key: u64) -> Result<Vec<u64>, HeaderError> {
1174 let len = self
1175 .d
1176 .array()
1177 .map_err(|_| HeaderError::Malformed("expected an array"))?
1178 .ok_or(HeaderError::Indefinite)?;
1179 if len > limits::MAX_LIST_ITEMS as u64 {
1180 return Err(HeaderError::ListTooLong {
1181 key,
1182 len,
1183 max: limits::MAX_LIST_ITEMS,
1184 });
1185 }
1186 let mut out = Vec::with_capacity(len as usize);
1188 for _ in 0..len {
1189 out.push(self.u64()?);
1190 }
1191 Ok(out)
1192 }
1193
1194 fn report_levels(&mut self) -> Result<Vec<CursorLevel>, HeaderError> {
1202 let len = self
1203 .d
1204 .array()
1205 .map_err(|_| HeaderError::Malformed("expected an array"))?
1206 .ok_or(HeaderError::Indefinite)?;
1207 if len > limits::MAX_REPORT_LEVELS as u64 {
1208 return Err(HeaderError::InvalidReport("too many report levels"));
1209 }
1210 let mut out: Vec<CursorLevel> = Vec::with_capacity(len as usize);
1212 let mut last: Option<u64> = None;
1213 for _ in 0..len {
1214 let value = self.u64()?;
1215 if let Some(prev) = last
1216 && value <= prev
1217 {
1218 return Err(HeaderError::InvalidReport("report levels must ascend"));
1219 }
1220 last = Some(value);
1221 out.push(
1222 CursorLevel::from_wire(value).ok_or(HeaderError::UnknownLevel {
1223 dimension: "report",
1224 value,
1225 })?,
1226 );
1227 }
1228 Ok(out)
1229 }
1230
1231 fn skip(&mut self) -> Result<(), HeaderError> {
1232 skip_value(self.d, limits::MAX_SKIP_DEPTH)
1233 }
1234}
1235
1236fn finish(d: &Decoder<'_>) -> Result<(), HeaderError> {
1238 if d.position() == d.input().len() {
1239 Ok(())
1240 } else {
1241 Err(HeaderError::TrailingBytes)
1242 }
1243}
1244
1245#[derive(Clone, Debug, PartialEq, Eq)]
1247pub struct Hello {
1248 pub versions: Vec<u64>,
1250 pub max_header_bytes: u64,
1252 pub max_transfers: u64,
1254 pub capabilities: Vec<u64>,
1256 pub required_capabilities: Vec<u64>,
1258 pub guarantees_offered: Option<GuaranteeSet>,
1264 pub guarantees_required: Option<GuaranteeSet>,
1270}
1271
1272impl Hello {
1273 pub fn v0(max_header_bytes: u64, max_transfers: u64) -> Hello {
1276 Hello {
1277 versions: vec![crate::VERSION],
1278 max_header_bytes,
1279 max_transfers,
1280 capabilities: Vec::new(),
1281 required_capabilities: Vec::new(),
1282 guarantees_offered: None,
1283 guarantees_required: None,
1284 }
1285 }
1286
1287 pub fn offered(&self) -> GuaranteeSet {
1289 self.guarantees_offered.unwrap_or(GuaranteeSet::CORE)
1290 }
1291
1292 pub fn required(&self) -> GuaranteeSet {
1294 self.guarantees_required.unwrap_or(GuaranteeSet::CORE)
1295 }
1296
1297 pub fn encode(&self) -> Vec<u8> {
1299 let mut out = Vec::new();
1300 self.encode_into(&mut out);
1301 out
1302 }
1303
1304 pub fn encode_into(&self, out: &mut Vec<u8>) {
1308 encode_into_with(out, |e| {
1309 let offered = self.guarantees_offered.filter(|s| !s.is_core());
1312 let required = self.guarantees_required.filter(|s| !s.is_core());
1313 e.map(5 + u64::from(offered.is_some()) + u64::from(required.is_some()))?;
1314 e.u64(hello_key::VERSIONS)?
1315 .array(self.versions.len() as u64)?;
1316 for v in &self.versions {
1317 e.u64(*v)?;
1318 }
1319 e.u64(hello_key::MAX_HEADER_BYTES)?
1320 .u64(self.max_header_bytes)?;
1321 e.u64(hello_key::MAX_TRANSFERS)?.u64(self.max_transfers)?;
1322 e.u64(hello_key::CAPABILITIES)?
1323 .array(self.capabilities.len() as u64)?;
1324 for c in &self.capabilities {
1325 e.u64(*c)?;
1326 }
1327 e.u64(hello_key::REQUIRED_CAPABILITIES)?
1328 .array(self.required_capabilities.len() as u64)?;
1329 for c in &self.required_capabilities {
1330 e.u64(*c)?;
1331 }
1332 if let Some(set) = offered {
1333 e.u64(hello_key::GUARANTEES_OFFERED)?;
1334 set.encode_into(e)?;
1335 }
1336 if let Some(set) = required {
1337 e.u64(hello_key::GUARANTEES_REQUIRED)?;
1338 set.encode_into(e)?;
1339 }
1340 Ok(())
1341 })
1342 }
1343
1344 pub fn decode(bytes: &[u8]) -> Result<Hello, HeaderError> {
1346 let mut d = Decoder::new(bytes);
1347 let mut versions = Vec::new();
1348 let mut max_header_bytes = 0;
1349 let mut max_transfers = 0;
1350 let mut capabilities = Vec::new();
1351 let mut required_capabilities = Vec::new();
1352 let mut guarantees_offered = None;
1353 let mut guarantees_required = None;
1354 {
1355 let mut m = MapReader::new(&mut d)?;
1356 while let Some(key) = m.next_key()? {
1357 match key {
1358 hello_key::VERSIONS => versions = m.uint_list(key)?,
1359 hello_key::MAX_HEADER_BYTES => max_header_bytes = m.u64()?,
1360 hello_key::MAX_TRANSFERS => max_transfers = m.u64()?,
1361 hello_key::CAPABILITIES => capabilities = m.uint_list(key)?,
1362 hello_key::REQUIRED_CAPABILITIES => required_capabilities = m.uint_list(key)?,
1363 hello_key::GUARANTEES_OFFERED => {
1364 guarantees_offered = Some(GuaranteeSet::decode_from(&mut m)?);
1365 }
1366 hello_key::GUARANTEES_REQUIRED => {
1367 guarantees_required = Some(GuaranteeSet::decode_from(&mut m)?);
1368 }
1369 _ => m.skip()?,
1370 }
1371 }
1372 for key in [
1373 hello_key::VERSIONS,
1374 hello_key::MAX_HEADER_BYTES,
1375 hello_key::MAX_TRANSFERS,
1376 hello_key::CAPABILITIES,
1377 hello_key::REQUIRED_CAPABILITIES,
1378 ] {
1379 m.require(key)?;
1380 }
1381 }
1382 finish(&d)?;
1383 let hello = Hello {
1384 versions,
1385 max_header_bytes,
1386 max_transfers,
1387 capabilities,
1388 required_capabilities,
1389 guarantees_offered,
1390 guarantees_required,
1391 };
1392 if !hello.offered().reaches(&hello.required()) {
1395 return Err(HeaderError::InvalidGuarantees(
1396 "guarantees_required is not covered by guarantees_offered",
1397 ));
1398 }
1399 Ok(hello)
1400 }
1401}
1402
1403#[derive(Clone, Debug, Default, PartialEq, Eq)]
1410pub struct DataHeader {
1411 pub endpoint: Option<String>,
1413 pub content_len: Option<u64>,
1415 pub content_type: Option<String>,
1417 pub traceparent: Option<String>,
1419 pub tracestate: Option<String>,
1421 pub topic: Option<String>,
1425 pub sequence: Option<u64>,
1434 pub producer: Option<[u8; limits::PRODUCER_BYTES]>,
1445 pub achieved: Option<Acknowledgement>,
1467 pub report_id: Option<u64>,
1474 pub report: Vec<CursorLevel>,
1483 pub report_mode: ReportMode,
1487 pub segment: Option<u64>,
1496 pub layer: Option<u8>,
1502}
1503
1504impl DataHeader {
1505 pub fn addressed(endpoint: impl Into<String>) -> DataHeader {
1507 DataHeader {
1508 endpoint: Some(endpoint.into()),
1509 ..DataHeader::default()
1510 }
1511 }
1512
1513 pub fn reply() -> DataHeader {
1518 DataHeader::default()
1519 }
1520
1521 pub fn encode(&self) -> Vec<u8> {
1538 let mut out = Vec::new();
1539 self.encode_into(&mut out);
1540 out
1541 }
1542
1543 pub fn encode_into(&self, out: &mut Vec<u8>) {
1547 let mut report: Vec<u64> = self.report.iter().map(|level| level.to_wire()).collect();
1550 report.sort_unstable();
1551 report.dedup();
1552 let count = u64::from(self.endpoint.is_some())
1553 + u64::from(self.content_len.is_some())
1554 + u64::from(self.content_type.is_some())
1555 + u64::from(self.traceparent.is_some())
1556 + u64::from(self.tracestate.is_some())
1557 + u64::from(self.topic.is_some())
1558 + u64::from(self.sequence.is_some())
1559 + u64::from(self.producer.is_some())
1560 + u64::from(self.achieved.is_some())
1561 + u64::from(self.report_id.is_some())
1562 + u64::from(!report.is_empty())
1563 + u64::from(self.report_mode != ReportMode::default())
1564 + u64::from(self.segment.is_some())
1565 + u64::from(self.layer.is_some());
1566 encode_into_with(out, |e| {
1567 e.map(count)?;
1568 if let Some(endpoint) = &self.endpoint {
1569 e.u64(data_key::ENDPOINT)?.str(endpoint)?;
1570 }
1571 if let Some(len) = self.content_len {
1572 e.u64(data_key::CONTENT_LEN)?.u64(len)?;
1573 }
1574 if let Some(ct) = &self.content_type {
1575 e.u64(data_key::CONTENT_TYPE)?.str(ct)?;
1576 }
1577 if let Some(tp) = &self.traceparent {
1578 e.u64(data_key::TRACEPARENT)?.str(tp)?;
1579 }
1580 if let Some(ts) = &self.tracestate {
1581 e.u64(data_key::TRACESTATE)?.str(ts)?;
1582 }
1583 if let Some(topic) = &self.topic {
1584 e.u64(data_key::TOPIC)?.str(topic)?;
1585 }
1586 if let Some(sequence) = self.sequence {
1589 e.u64(data_key::SEQUENCE)?.u64(sequence)?;
1590 }
1591 if let Some(producer) = &self.producer {
1592 e.u64(data_key::PRODUCER)?.bytes(producer)?;
1593 }
1594 if let Some(achieved) = self.achieved {
1595 e.u64(data_key::ACHIEVED)?.u64(achieved.to_wire())?;
1596 }
1597 if let Some(report_id) = self.report_id {
1598 e.u64(data_key::REPORT_ID)?.u64(report_id)?;
1599 }
1600 if !report.is_empty() {
1601 e.u64(data_key::REPORT)?.array(report.len() as u64)?;
1602 for value in &report {
1603 e.u64(*value)?;
1604 }
1605 }
1606 if self.report_mode != ReportMode::default() {
1610 e.u64(data_key::REPORT_MODE)?
1611 .u64(self.report_mode.to_wire())?;
1612 }
1613 if let Some(segment) = self.segment {
1614 e.u64(data_key::SEGMENT)?.u64(segment)?;
1615 }
1616 if let Some(layer) = self.layer {
1617 e.u64(data_key::LAYER)?.u64(u64::from(layer))?;
1618 }
1619 Ok(())
1620 })
1621 }
1622
1623 pub fn decode(bytes: &[u8]) -> Result<DataHeader, HeaderError> {
1625 let mut d = Decoder::new(bytes);
1626 let mut header = DataHeader::default();
1627 {
1628 let mut m = MapReader::new(&mut d)?;
1629 while let Some(key) = m.next_key()? {
1630 match key {
1631 data_key::ENDPOINT => {
1632 header.endpoint = Some(m.text(key, limits::MAX_ENDPOINT_BYTES)?)
1633 }
1634 data_key::CONTENT_LEN => header.content_len = Some(m.u64()?),
1635 data_key::CONTENT_TYPE => {
1636 header.content_type = Some(m.text(key, limits::MAX_CONTENT_TYPE_BYTES)?)
1637 }
1638 data_key::TRACEPARENT => {
1639 header.traceparent = Some(m.text(key, limits::MAX_TRACEPARENT_BYTES)?)
1640 }
1641 data_key::TRACESTATE => {
1642 header.tracestate = Some(m.text(key, limits::MAX_TRACESTATE_BYTES)?)
1643 }
1644 data_key::TOPIC => header.topic = Some(m.text(key, limits::MAX_TOPIC_BYTES)?),
1645 data_key::SEQUENCE => header.sequence = Some(m.u64()?),
1646 data_key::PRODUCER => header.producer = Some(m.byte_array(key)?),
1647 data_key::ACHIEVED => {
1652 let value = m.u64()?;
1653 header.achieved = Some(Acknowledgement::from_wire(value).ok_or(
1654 HeaderError::UnknownLevel {
1655 dimension: "achieved",
1656 value,
1657 },
1658 )?);
1659 }
1660 data_key::REPORT_ID => header.report_id = Some(m.u64()?),
1661 data_key::REPORT => header.report = m.report_levels()?,
1662 data_key::REPORT_MODE => {
1663 header.report_mode = level(m.u64()?, "report_mode")?;
1664 }
1665 data_key::SEGMENT => header.segment = Some(m.u64()?),
1666 data_key::LAYER => header.layer = Some(layer(m.u64()?, "layer above 15")?),
1667 _ => m.skip()?,
1668 }
1669 }
1670 }
1671 finish(&d)?;
1672 if header.layer.is_some() && header.segment.is_none() {
1675 return Err(HeaderError::InvalidLayer("layer without segment"));
1676 }
1677 if !header.report.is_empty() && header.report_id.is_none() {
1681 return Err(HeaderError::InvalidReport("report without report_id"));
1682 }
1683 if header.report_id.is_some() && header.report.is_empty() {
1684 return Err(HeaderError::InvalidReport("report_id without report"));
1685 }
1686 Ok(header)
1687 }
1688}
1689
1690#[derive(Clone, Debug, PartialEq, Eq)]
1695pub struct ErrorHeader {
1696 pub code: u64,
1698 pub message: Option<String>,
1700}
1701
1702impl ErrorHeader {
1703 pub fn new(code: weida_core::ErrorCode) -> ErrorHeader {
1705 ErrorHeader {
1706 code: code.to_wire(),
1707 message: None,
1708 }
1709 }
1710
1711 pub fn error_code(&self) -> Option<weida_core::ErrorCode> {
1713 weida_core::ErrorCode::from_wire(self.code)
1714 }
1715
1716 pub fn encode(&self) -> Vec<u8> {
1718 let mut out = Vec::new();
1719 self.encode_into(&mut out);
1720 out
1721 }
1722
1723 pub fn encode_into(&self, out: &mut Vec<u8>) {
1727 let count = 1 + u64::from(self.message.is_some());
1728 encode_into_with(out, |e| {
1729 e.map(count)?;
1730 e.u64(error_key::CODE)?.u64(self.code)?;
1731 if let Some(msg) = &self.message {
1732 e.u64(error_key::MESSAGE)?.str(msg)?;
1733 }
1734 Ok(())
1735 })
1736 }
1737
1738 pub fn decode(bytes: &[u8]) -> Result<ErrorHeader, HeaderError> {
1740 let mut d = Decoder::new(bytes);
1741 let mut code = 0;
1742 let mut message = None;
1743 {
1744 let mut m = MapReader::new(&mut d)?;
1745 while let Some(key) = m.next_key()? {
1746 match key {
1747 error_key::CODE => code = m.u64()?,
1748 error_key::MESSAGE => message = Some(m.text(key, limits::MAX_MESSAGE_BYTES)?),
1749 _ => m.skip()?,
1750 }
1751 }
1752 m.require(error_key::CODE)?;
1753 }
1754 finish(&d)?;
1755 Ok(ErrorHeader { code, message })
1756 }
1757}
1758
1759#[derive(Clone, Debug, PartialEq, Eq)]
1766pub struct SubscriptionHeader {
1767 pub endpoint: String,
1769 pub filter: String,
1772 pub max_age_ms: Option<u64>,
1777 pub max_layer: Option<u8>,
1782}
1783
1784impl SubscriptionHeader {
1785 pub fn new(endpoint: impl Into<String>, filter: impl Into<String>) -> SubscriptionHeader {
1787 SubscriptionHeader {
1788 endpoint: endpoint.into(),
1789 filter: filter.into(),
1790 max_age_ms: None,
1791 max_layer: None,
1792 }
1793 }
1794
1795 pub fn encode(&self) -> Vec<u8> {
1797 let mut out = Vec::new();
1798 self.encode_into(&mut out);
1799 out
1800 }
1801
1802 pub fn encode_into(&self, out: &mut Vec<u8>) {
1806 encode_into_with(out, |e| {
1811 e.map(2 + u64::from(self.max_age_ms.is_some()) + u64::from(self.max_layer.is_some()))?;
1812 e.u64(subscription_key::ENDPOINT)?.str(&self.endpoint)?;
1813 e.u64(subscription_key::FILTER)?.str(&self.filter)?;
1814 if let Some(max_age_ms) = self.max_age_ms {
1815 e.u64(subscription_key::MAX_AGE_MS)?.u64(max_age_ms)?;
1816 }
1817 if let Some(max_layer) = self.max_layer {
1818 e.u64(subscription_key::MAX_LAYER)?
1819 .u64(u64::from(max_layer))?;
1820 }
1821 Ok(())
1822 })
1823 }
1824
1825 pub fn decode(bytes: &[u8]) -> Result<SubscriptionHeader, HeaderError> {
1827 let mut d = Decoder::new(bytes);
1828 let mut endpoint = None;
1829 let mut filter = None;
1830 let mut max_age_ms = None;
1831 let mut max_layer = None;
1832 {
1833 let mut m = MapReader::new(&mut d)?;
1834 while let Some(key) = m.next_key()? {
1835 match key {
1836 subscription_key::ENDPOINT => {
1837 endpoint = Some(m.text(key, limits::MAX_ENDPOINT_BYTES)?)
1838 }
1839 subscription_key::FILTER => {
1840 filter = Some(m.text(key, limits::MAX_FILTER_BYTES)?)
1841 }
1842 subscription_key::MAX_AGE_MS => max_age_ms = Some(m.u64()?),
1843 subscription_key::MAX_LAYER => {
1844 max_layer = Some(layer(m.u64()?, "max_layer above 15")?);
1845 }
1846 _ => m.skip()?,
1847 }
1848 }
1849 m.require(subscription_key::ENDPOINT)?;
1850 m.require(subscription_key::FILTER)?;
1851 }
1852 let filter = filter.expect("presence checked above");
1857 filter::validate(&filter)?;
1858 finish(&d)?;
1859 Ok(SubscriptionHeader {
1860 endpoint: endpoint.expect("presence checked above"),
1861 filter,
1862 max_age_ms,
1863 max_layer,
1864 })
1865 }
1866}
1867
1868#[derive(Clone, Debug, PartialEq, Eq)]
1875pub struct FlowHeader {
1876 pub endpoint: String,
1878 pub flow: u64,
1881 pub content_type: Option<String>,
1883 pub traceparent: Option<String>,
1885 pub tracestate: Option<String>,
1887 pub topic: Option<String>,
1889}
1890
1891impl FlowHeader {
1892 pub fn new(endpoint: impl Into<String>, flow: u64) -> FlowHeader {
1894 FlowHeader {
1895 endpoint: endpoint.into(),
1896 flow,
1897 content_type: None,
1898 traceparent: None,
1899 tracestate: None,
1900 topic: None,
1901 }
1902 }
1903
1904 pub fn encode(&self) -> Vec<u8> {
1906 let mut out = Vec::new();
1907 self.encode_into(&mut out);
1908 out
1909 }
1910
1911 pub fn encode_into(&self, out: &mut Vec<u8>) {
1914 let count = 2
1915 + u64::from(self.content_type.is_some())
1916 + u64::from(self.traceparent.is_some())
1917 + u64::from(self.tracestate.is_some())
1918 + u64::from(self.topic.is_some());
1919 encode_into_with(out, |e| {
1920 e.map(count)?;
1921 e.u64(flow_key::ENDPOINT)?.str(&self.endpoint)?;
1922 e.u64(flow_key::FLOW)?.u64(self.flow)?;
1923 if let Some(ct) = &self.content_type {
1924 e.u64(flow_key::CONTENT_TYPE)?.str(ct)?;
1925 }
1926 if let Some(tp) = &self.traceparent {
1927 e.u64(flow_key::TRACEPARENT)?.str(tp)?;
1928 }
1929 if let Some(ts) = &self.tracestate {
1930 e.u64(flow_key::TRACESTATE)?.str(ts)?;
1931 }
1932 if let Some(topic) = &self.topic {
1933 e.u64(flow_key::TOPIC)?.str(topic)?;
1934 }
1935 Ok(())
1936 })
1937 }
1938
1939 pub fn decode(bytes: &[u8]) -> Result<FlowHeader, HeaderError> {
1941 let mut d = Decoder::new(bytes);
1942 let mut endpoint = None;
1943 let mut flow = None;
1944 let mut content_type = None;
1945 let mut traceparent = None;
1946 let mut tracestate = None;
1947 let mut topic = None;
1948 {
1949 let mut m = MapReader::new(&mut d)?;
1950 while let Some(key) = m.next_key()? {
1951 match key {
1952 flow_key::ENDPOINT => endpoint = Some(m.text(key, limits::MAX_ENDPOINT_BYTES)?),
1953 flow_key::FLOW => flow = Some(m.u64()?),
1954 flow_key::CONTENT_TYPE => {
1955 content_type = Some(m.text(key, limits::MAX_CONTENT_TYPE_BYTES)?)
1956 }
1957 flow_key::TRACEPARENT => {
1958 traceparent = Some(m.text(key, limits::MAX_TRACEPARENT_BYTES)?)
1959 }
1960 flow_key::TRACESTATE => {
1961 tracestate = Some(m.text(key, limits::MAX_TRACESTATE_BYTES)?)
1962 }
1963 flow_key::TOPIC => topic = Some(m.text(key, limits::MAX_TOPIC_BYTES)?),
1964 _ => m.skip()?,
1965 }
1966 }
1967 m.require(flow_key::ENDPOINT)?;
1968 m.require(flow_key::FLOW)?;
1969 }
1970 finish(&d)?;
1971 Ok(FlowHeader {
1972 endpoint: endpoint.expect("presence checked above"),
1973 flow: flow.expect("presence checked above"),
1974 content_type,
1975 traceparent,
1976 tracestate,
1977 topic,
1978 })
1979 }
1980}
1981
1982#[derive(Clone, Debug, PartialEq, Eq)]
2006pub struct CreditHeader {
2007 pub endpoint: String,
2009 pub filter: String,
2012 pub limit: u64,
2015}
2016
2017impl CreditHeader {
2018 pub fn new(endpoint: impl Into<String>, filter: impl Into<String>, limit: u64) -> CreditHeader {
2020 CreditHeader {
2021 endpoint: endpoint.into(),
2022 filter: filter.into(),
2023 limit,
2024 }
2025 }
2026
2027 pub fn encode(&self) -> Vec<u8> {
2029 let mut out = Vec::new();
2030 self.encode_into(&mut out);
2031 out
2032 }
2033
2034 pub fn encode_into(&self, out: &mut Vec<u8>) {
2038 encode_into_with(out, |e| {
2039 e.map(3)?;
2040 e.u64(credit_key::ENDPOINT)?.str(&self.endpoint)?;
2041 e.u64(credit_key::FILTER)?.str(&self.filter)?;
2042 e.u64(credit_key::LIMIT)?.u64(self.limit)?;
2043 Ok(())
2044 })
2045 }
2046
2047 pub fn decode(bytes: &[u8]) -> Result<CreditHeader, HeaderError> {
2049 let mut d = Decoder::new(bytes);
2050 let mut endpoint = None;
2051 let mut filter = None;
2052 let mut limit = None;
2053 {
2054 let mut m = MapReader::new(&mut d)?;
2055 while let Some(key) = m.next_key()? {
2056 match key {
2057 credit_key::ENDPOINT => {
2058 endpoint = Some(m.text(key, limits::MAX_ENDPOINT_BYTES)?)
2059 }
2060 credit_key::FILTER => filter = Some(m.text(key, limits::MAX_FILTER_BYTES)?),
2061 credit_key::LIMIT => limit = Some(m.u64()?),
2062 _ => m.skip()?,
2063 }
2064 }
2065 m.require(credit_key::ENDPOINT)?;
2066 m.require(credit_key::FILTER)?;
2067 m.require(credit_key::LIMIT)?;
2068 }
2069 let filter = filter.expect("presence checked above");
2072 filter::validate(&filter)?;
2073 finish(&d)?;
2074 Ok(CreditHeader {
2075 endpoint: endpoint.expect("presence checked above"),
2076 filter,
2077 limit: limit.expect("presence checked above"),
2078 })
2079 }
2080}
2081
2082#[derive(Clone, Copy, Debug, PartialEq, Eq)]
2090pub struct CursorHeader {
2091 pub report_id: u64,
2093}
2094
2095impl CursorHeader {
2096 pub fn encode(&self) -> Vec<u8> {
2098 let mut out = Vec::new();
2099 self.encode_into(&mut out);
2100 out
2101 }
2102
2103 pub fn encode_into(&self, out: &mut Vec<u8>) {
2107 encode_into_with(out, |e| {
2108 e.map(1)?;
2109 e.u64(cursor_key::REPORT_ID)?.u64(self.report_id)?;
2110 Ok(())
2111 })
2112 }
2113
2114 pub fn decode(bytes: &[u8]) -> Result<CursorHeader, HeaderError> {
2116 let mut d = Decoder::new(bytes);
2117 let mut report_id = None;
2118 {
2119 let mut m = MapReader::new(&mut d)?;
2120 while let Some(key) = m.next_key()? {
2121 match key {
2122 cursor_key::REPORT_ID => report_id = Some(m.u64()?),
2123 _ => m.skip()?,
2124 }
2125 }
2126 m.require(cursor_key::REPORT_ID)?;
2127 }
2128 finish(&d)?;
2129 Ok(CursorHeader {
2130 report_id: report_id.expect("presence checked above"),
2131 })
2132 }
2133}
2134
2135pub const MAX_CURSOR_RECORD_LEN: usize = 2 * crate::varint::MAX_ENCODED_LEN;
2137
2138pub fn encode_cursor_record(
2145 level: CursorLevel,
2146 offset: u64,
2147 out: &mut Vec<u8>,
2148) -> Result<(), VarintError> {
2149 encode_varint(level.to_wire(), out)?;
2150 encode_varint(offset, out)
2151}
2152
2153pub fn decode_cursor_record(
2162 input: &[u8],
2163) -> Result<Option<(CursorLevel, u64, usize)>, HeaderError> {
2164 let Ok((raw_level, level_len)) = decode_varint(input) else {
2167 return Ok(None);
2168 };
2169 let Ok((offset, offset_len)) = decode_varint(&input[level_len..]) else {
2170 return Ok(None);
2171 };
2172 let level = CursorLevel::from_wire(raw_level).ok_or(HeaderError::UnknownLevel {
2173 dimension: "cursor",
2174 value: raw_level,
2175 })?;
2176 Ok(Some((level, offset, level_len + offset_len)))
2177}
2178
2179#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
2182pub struct ReportHead;
2183
2184impl ReportHead {
2185 pub fn encode(&self) -> Vec<u8> {
2187 let mut out = Vec::new();
2188 encode_into_with(&mut out, |e| {
2189 e.map(0)?;
2190 Ok(())
2191 });
2192 out
2193 }
2194
2195 pub fn decode(bytes: &[u8]) -> Result<ReportHead, HeaderError> {
2197 let mut d = Decoder::new(bytes);
2198 {
2199 let mut m = MapReader::new(&mut d)?;
2200 while m.next_key()?.is_some() {
2201 m.skip()?;
2202 }
2203 }
2204 finish(&d)?;
2205 Ok(ReportHead)
2206 }
2207}
2208
2209mod path_key {
2211 pub const RTT_US: u64 = 0;
2212 pub const MIN_RTT_US: u64 = 1;
2213 pub const CWND: u64 = 2;
2214 pub const CONGESTION_EVENTS: u64 = 3;
2215 pub const LOST_PACKETS: u64 = 4;
2216 pub const LOST_BYTES: u64 = 5;
2217 pub const SENT_PACKETS: u64 = 6;
2218 pub const CURRENT_MTU: u64 = 7;
2219 pub const TX_DATAGRAMS: u64 = 8;
2220 pub const TX_BYTES: u64 = 9;
2221 pub const RX_DATAGRAMS: u64 = 10;
2222 pub const RX_BYTES: u64 = 11;
2223}
2224
2225#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
2231pub struct PathRecord {
2232 pub rtt_us: u64,
2234 pub min_rtt_us: u64,
2236 pub cwnd: u64,
2238 pub congestion_events: u64,
2240 pub lost_packets: u64,
2242 pub lost_bytes: u64,
2244 pub sent_packets: u64,
2246 pub current_mtu: u64,
2248 pub tx_datagrams: u64,
2250 pub tx_bytes: u64,
2252 pub rx_datagrams: u64,
2254 pub rx_bytes: u64,
2256}
2257
2258impl PathRecord {
2259 fn fields(&self) -> [(u64, u64); 12] {
2260 [
2261 (path_key::RTT_US, self.rtt_us),
2262 (path_key::MIN_RTT_US, self.min_rtt_us),
2263 (path_key::CWND, self.cwnd),
2264 (path_key::CONGESTION_EVENTS, self.congestion_events),
2265 (path_key::LOST_PACKETS, self.lost_packets),
2266 (path_key::LOST_BYTES, self.lost_bytes),
2267 (path_key::SENT_PACKETS, self.sent_packets),
2268 (path_key::CURRENT_MTU, self.current_mtu),
2269 (path_key::TX_DATAGRAMS, self.tx_datagrams),
2270 (path_key::TX_BYTES, self.tx_bytes),
2271 (path_key::RX_DATAGRAMS, self.rx_datagrams),
2272 (path_key::RX_BYTES, self.rx_bytes),
2273 ]
2274 }
2275
2276 pub fn encode_into(&self, out: &mut Vec<u8>) {
2279 let fields = self.fields();
2280 let present = fields.iter().filter(|(_, v)| *v != 0).count() as u64;
2281 let mut map = Vec::new();
2282 encode_into_with(&mut map, |e| {
2283 e.map(present)?;
2284 for (key, value) in fields {
2285 if value != 0 {
2286 e.u64(key)?.u64(value)?;
2287 }
2288 }
2289 Ok(())
2290 });
2291 encode_varint(map.len() as u64, out).expect("a record is far below 2^62 bytes");
2292 out.extend_from_slice(&map);
2293 }
2294
2295 pub fn decode(input: &[u8]) -> Result<Option<(PathRecord, usize)>, HeaderError> {
2303 let Ok((len, prefix)) = decode_varint(input) else {
2304 return Ok(None);
2305 };
2306 if len > limits::MAX_PATH_RECORD_BYTES as u64 {
2307 return Err(HeaderError::InvalidPathReport("record above 256 bytes"));
2308 }
2309 let len = len as usize;
2310 let Some(map) = input.get(prefix..prefix + len) else {
2311 return Ok(None);
2312 };
2313 let mut record = PathRecord::default();
2314 let mut d = Decoder::new(map);
2315 {
2316 let mut m = MapReader::new(&mut d)?;
2317 while let Some(key) = m.next_key()? {
2318 let slot = match key {
2319 path_key::RTT_US => &mut record.rtt_us,
2320 path_key::MIN_RTT_US => &mut record.min_rtt_us,
2321 path_key::CWND => &mut record.cwnd,
2322 path_key::CONGESTION_EVENTS => &mut record.congestion_events,
2323 path_key::LOST_PACKETS => &mut record.lost_packets,
2324 path_key::LOST_BYTES => &mut record.lost_bytes,
2325 path_key::SENT_PACKETS => &mut record.sent_packets,
2326 path_key::CURRENT_MTU => &mut record.current_mtu,
2327 path_key::TX_DATAGRAMS => &mut record.tx_datagrams,
2328 path_key::TX_BYTES => &mut record.tx_bytes,
2329 path_key::RX_DATAGRAMS => &mut record.rx_datagrams,
2330 path_key::RX_BYTES => &mut record.rx_bytes,
2331 _ => {
2332 m.skip()?;
2333 continue;
2334 }
2335 };
2336 *slot = m.u64()?;
2337 }
2338 }
2339 finish(&d)?;
2340 Ok(Some((record, prefix + len)))
2341 }
2342}
2343
2344#[cfg(test)]
2345mod tests {
2346 fn encode_with(
2349 f: impl FnOnce(
2350 &mut minicbor::Encoder<Vec<u8>>,
2351 ) -> Result<(), minicbor::encode::Error<std::convert::Infallible>>,
2352 ) -> Vec<u8> {
2353 let mut out = Vec::new();
2354 super::encode_into_with(&mut out, f);
2355 out
2356 }
2357
2358 use super::*;
2359 use weida_core::ErrorCode;
2360
2361 #[test]
2364 fn golden_data_request_header() {
2365 let h = DataHeader::addressed("/t");
2366 let bytes = h.encode();
2367 assert_eq!(bytes, vec![0xA1, 0x00, 0x62, 0x2F, 0x74]);
2368 assert_eq!(bytes.len(), 0x05);
2369 assert_eq!(DataHeader::decode(&bytes).unwrap(), h);
2370 }
2371
2372 #[test]
2373 fn golden_data_reply_header() {
2374 let h = DataHeader::reply();
2376 let bytes = h.encode();
2377 assert_eq!(bytes, vec![0xA0]);
2378 assert_eq!(bytes.len(), 0x01);
2379 assert_eq!(DataHeader::decode(&bytes).unwrap(), h);
2380 }
2381
2382 #[test]
2383 fn golden_hello_header() {
2384 let h = Hello::v0(16384, 1024);
2385 let bytes = h.encode();
2386 assert_eq!(
2387 bytes,
2388 vec![
2389 0xA5, 0x00, 0x81, 0x00, 0x01, 0x19, 0x40, 0x00, 0x02, 0x19, 0x04, 0x00, 0x03, 0x80,
2390 0x04, 0x80
2391 ]
2392 );
2393 assert_eq!(bytes.len(), 0x10);
2394 assert_eq!(Hello::decode(&bytes).unwrap(), h);
2395 }
2396
2397 #[test]
2398 fn golden_error_header() {
2399 let h = ErrorHeader::new(ErrorCode::NoReply);
2400 let bytes = h.encode();
2401 assert_eq!(bytes, vec![0xA1, 0x00, 0x05]);
2402 assert_eq!(bytes.len(), 0x03);
2403 assert_eq!(ErrorHeader::decode(&bytes).unwrap(), h);
2404 }
2405
2406 #[test]
2407 fn golden_pub_copy_data_header() {
2408 let mut h = DataHeader::addressed("/md");
2409 h.topic = Some("px.eur".into());
2410 let bytes = h.encode();
2411 assert_eq!(
2412 bytes,
2413 vec![
2414 0xA2, 0x00, 0x63, 0x2F, 0x6D, 0x64, 0x05, 0x66, 0x70, 0x78, 0x2E, 0x65, 0x75, 0x72
2415 ]
2416 );
2417 assert_eq!(bytes.len(), 0x0E);
2418 assert_eq!(DataHeader::decode(&bytes).unwrap(), h);
2419 }
2420
2421 const VECTOR_PRODUCER: [u8; limits::PRODUCER_BYTES] = [
2424 0x9F, 0x86, 0xD0, 0x81, 0x88, 0x4C, 0x7D, 0x65, 0x9A, 0x2F, 0xEA, 0xA0, 0xC5, 0x5A, 0xD0,
2425 0x15, 0xA3, 0xBF, 0x4F, 0x1B, 0x2B, 0x0B, 0x82, 0x2C, 0xD1, 0x5D, 0x6C, 0x15, 0xB0, 0xF0,
2426 0x0A, 0x08,
2427 ];
2428
2429 #[test]
2430 fn golden_sequenced_data_header() {
2431 let mut h = DataHeader::addressed("/t");
2432 h.sequence = Some(1);
2433 let bytes = h.encode();
2434 assert_eq!(bytes, vec![0xA2, 0x00, 0x62, 0x2F, 0x74, 0x06, 0x01]);
2435 assert_eq!(bytes.len(), 0x07);
2436 assert_eq!(DataHeader::decode(&bytes).unwrap(), h);
2437 }
2438
2439 #[test]
2440 fn golden_relayed_data_header() {
2441 let mut h = DataHeader::addressed("/t");
2442 h.sequence = Some(1);
2443 h.producer = Some(VECTOR_PRODUCER);
2444 let bytes = h.encode();
2445 let mut expected = vec![0xA3, 0x00, 0x62, 0x2F, 0x74, 0x06, 0x01, 0x07, 0x58, 0x20];
2446 expected.extend_from_slice(&VECTOR_PRODUCER);
2447 assert_eq!(bytes, expected);
2448 assert_eq!(bytes.len(), 0x2A);
2449 assert_eq!(DataHeader::decode(&bytes).unwrap(), h);
2450 }
2451
2452 #[test]
2453 fn a_producer_longer_than_the_cap_is_rejected() {
2454 let bytes = encode_with(|e| {
2455 e.map(1)?;
2456 e.u64(data_key::PRODUCER)?
2457 .bytes(&[0u8; limits::PRODUCER_BYTES + 1])?;
2458 Ok(())
2459 });
2460 assert_eq!(
2461 DataHeader::decode(&bytes),
2462 Err(HeaderError::StringTooLong {
2463 key: data_key::PRODUCER,
2464 len: limits::PRODUCER_BYTES + 1,
2465 max: limits::PRODUCER_BYTES,
2466 })
2467 );
2468 }
2469
2470 #[test]
2471 fn a_producer_shorter_than_a_digest_is_rejected() {
2472 let bytes = encode_with(|e| {
2475 e.map(1)?;
2476 e.u64(data_key::PRODUCER)?.bytes(&[0u8; 16])?;
2477 Ok(())
2478 });
2479 assert!(matches!(
2480 DataHeader::decode(&bytes),
2481 Err(HeaderError::Malformed(_))
2482 ));
2483 }
2484
2485 #[test]
2486 fn the_new_keys_reject_the_wrong_cbor_type() {
2487 let sequence_as_text = encode_with(|e| {
2488 e.map(1)?;
2489 e.u64(data_key::SEQUENCE)?.str("7")?;
2490 Ok(())
2491 });
2492 assert!(DataHeader::decode(&sequence_as_text).is_err());
2493
2494 let producer_as_text = encode_with(|e| {
2495 e.map(1)?;
2496 e.u64(data_key::PRODUCER)?.str("sha256:…")?;
2497 Ok(())
2498 });
2499 assert!(DataHeader::decode(&producer_as_text).is_err());
2500 }
2501
2502 #[test]
2503 fn a_v0_header_carries_neither_new_key() {
2504 let mut h = DataHeader::addressed("/t");
2508 h.traceparent = Some("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01".into());
2509 let bytes = h.encode();
2510 let mut d = Decoder::new(&bytes);
2511 let pairs = d.map().unwrap().unwrap();
2512 let keys: Vec<u64> = (0..pairs)
2513 .map(|_| {
2514 let key = d.u64().unwrap();
2515 d.skip().unwrap();
2516 key
2517 })
2518 .collect();
2519 assert_eq!(keys, vec![data_key::ENDPOINT, data_key::TRACEPARENT]);
2520 }
2521
2522 #[test]
2523 fn golden_subscription_headers() {
2524 let h = SubscriptionHeader::new("/md", "px.");
2525 let bytes = h.encode();
2526 assert_eq!(
2527 bytes,
2528 vec![
2529 0xA2, 0x00, 0x63, 0x2F, 0x6D, 0x64, 0x01, 0x63, 0x70, 0x78, 0x2E
2530 ]
2531 );
2532 assert_eq!(bytes.len(), 0x0B);
2533 assert_eq!(SubscriptionHeader::decode(&bytes).unwrap(), h);
2536 }
2537
2538 #[test]
2541 fn data_header_roundtrip_with_every_field() {
2542 let h = DataHeader {
2543 endpoint: Some("/transform".into()),
2544 content_len: Some(1 << 40),
2545 content_type: Some("application/octet-stream".into()),
2546 traceparent: Some("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01".into()),
2547 tracestate: Some("vendor=value".into()),
2548 topic: Some("px.eur".into()),
2549 sequence: Some(u64::MAX),
2550 producer: Some([0x5A; limits::PRODUCER_BYTES]),
2551 achieved: Some(Acknowledgement::Processed),
2552 report_id: Some(7),
2553 report: vec![
2554 CursorLevel::Known(Acknowledgement::Accepted),
2555 CursorLevel::Known(Acknowledgement::Processed),
2556 CursorLevel::Application(CursorLevel::APPLICATION_FLOOR),
2557 ],
2558 report_mode: ReportMode::FinalOnly,
2559 segment: Some(u64::MAX),
2560 layer: Some(limits::MAX_LAYER),
2561 };
2562 assert_eq!(DataHeader::decode(&h.encode()).unwrap(), h);
2563 }
2564
2565 #[test]
2566 fn error_header_roundtrip_with_and_without_message() {
2567 let bare = ErrorHeader::new(ErrorCode::UnknownEndpoint);
2568 assert_eq!(ErrorHeader::decode(&bare.encode()).unwrap(), bare);
2569 assert_eq!(bare.error_code(), Some(ErrorCode::UnknownEndpoint));
2570
2571 let with_msg = ErrorHeader {
2572 code: 4,
2573 message: Some("handler panicked".into()),
2574 };
2575 assert_eq!(ErrorHeader::decode(&with_msg.encode()).unwrap(), with_msg);
2576 }
2577
2578 #[test]
2579 fn keys_are_emitted_in_ascending_order() {
2580 let h = DataHeader {
2581 endpoint: Some("/x".into()),
2582 content_len: Some(1),
2583 content_type: Some("t".into()),
2584 traceparent: Some("p".into()),
2585 tracestate: Some("s".into()),
2586 topic: Some("k".into()),
2587 sequence: Some(9),
2588 producer: Some([0u8; limits::PRODUCER_BYTES]),
2589 achieved: Some(Acknowledgement::Accepted),
2590 report_id: Some(1),
2591 report: vec![CursorLevel::Known(Acknowledgement::Stored)],
2592 report_mode: ReportMode::FinalOnly,
2593 segment: Some(3),
2594 layer: Some(2),
2595 };
2596 let bytes = h.encode();
2597 let mut d = Decoder::new(&bytes);
2598 let n = d.map().unwrap().unwrap();
2599 let mut last = None;
2600 for _ in 0..n {
2601 let key = d.u64().unwrap();
2602 if let Some(prev) = last {
2603 assert!(key > prev, "keys must ascend: {prev} then {key}");
2604 }
2605 last = Some(key);
2606 d.skip().unwrap();
2607 }
2608 }
2609
2610 #[test]
2611 fn an_unsorted_report_with_repeats_is_emitted_as_the_canonical_ascending_set() {
2612 let h = DataHeader {
2617 report_id: Some(1),
2618 report: vec![
2619 CursorLevel::Application(CursorLevel::APPLICATION_FLOOR),
2620 CursorLevel::Known(Acknowledgement::Stored),
2621 CursorLevel::Application(CursorLevel::APPLICATION_FLOOR),
2622 CursorLevel::Known(Acknowledgement::Accepted),
2623 CursorLevel::Known(Acknowledgement::Stored),
2624 ],
2625 ..DataHeader::reply()
2626 };
2627 let decoded = DataHeader::decode(&h.encode()).expect("encoder emits the canonical form");
2628 assert_eq!(
2629 decoded.report,
2630 vec![
2631 CursorLevel::Known(Acknowledgement::Accepted),
2632 CursorLevel::Known(Acknowledgement::Stored),
2633 CursorLevel::Application(CursorLevel::APPLICATION_FLOOR),
2634 ]
2635 );
2636 }
2637
2638 #[test]
2641 fn every_data_field_is_optional_at_the_decoder() {
2642 assert_eq!(DataHeader::decode(&[0xA0]).unwrap(), DataHeader::default());
2646
2647 let only_topic = encode_with(|e| {
2648 e.map(1)?;
2649 e.u64(data_key::TOPIC)?.str("px.eur")?;
2650 Ok(())
2651 });
2652 let h = DataHeader::decode(&only_topic).unwrap();
2653 assert_eq!(h.topic.as_deref(), Some("px.eur"));
2654 assert_eq!(h.endpoint, None);
2655 }
2656
2657 #[test]
2658 fn absent_fields_are_omitted_by_the_encoder() {
2659 let h = DataHeader::addressed("/t");
2660 assert_eq!(h.encode(), vec![0xA1, 0x00, 0x62, 0x2F, 0x74]);
2661 }
2662
2663 #[test]
2666 fn unknown_keys_are_skipped() {
2667 let h = DataHeader::addressed("/t");
2670 let extended = encode_with(|e| {
2671 e.map(2)?;
2672 e.u64(0)?.str("/t")?;
2673 e.u64(63)?.array(2)?.u64(7)?.map(1)?.u64(1)?.bool(true)?;
2674 Ok(())
2675 });
2676 assert_eq!(DataHeader::decode(&extended).unwrap(), h);
2677 }
2678
2679 #[test]
2680 fn unknown_keys_above_the_reserved_range_are_skipped() {
2681 let extended = encode_with(|e| {
2682 e.map(2)?;
2683 e.u64(1)?.u64(5)?;
2684 e.u64(1000)?.str("future")?;
2685 Ok(())
2686 });
2687 let h = DataHeader::decode(&extended).unwrap();
2688 assert_eq!(h.content_len, Some(5));
2689 }
2690
2691 #[test]
2692 fn skipping_tolerates_nesting_up_to_the_depth_limit() {
2693 for depth in [1usize, limits::MAX_SKIP_DEPTH] {
2694 let bytes = encode_with(|e| {
2695 e.map(2)?;
2696 e.u64(data_key::CONTENT_LEN)?.u64(1)?;
2697 e.u64(50)?;
2698 for _ in 0..depth {
2699 e.array(1)?;
2700 }
2701 e.u64(1)?;
2702 Ok(())
2703 });
2704 let h = DataHeader::decode(&bytes).unwrap_or_else(|e| panic!("depth {depth}: {e}"));
2705 assert_eq!(h.content_len, Some(1), "depth {depth}");
2706 }
2707 }
2708
2709 #[test]
2710 fn skipping_rejects_nesting_beyond_the_depth_limit() {
2711 let bytes = encode_with(|e| {
2712 e.map(1)?;
2713 e.u64(50)?;
2714 for _ in 0..(limits::MAX_SKIP_DEPTH + 1) {
2715 e.array(1)?;
2716 }
2717 e.u64(1)?;
2718 Ok(())
2719 });
2720 assert_eq!(
2721 DataHeader::decode(&bytes).unwrap_err(),
2722 HeaderError::DepthExceeded
2723 );
2724 }
2725
2726 #[test]
2727 fn skipping_a_wide_shallow_structure_is_fine() {
2728 let bytes = encode_with(|e| {
2729 e.map(2)?;
2730 e.u64(data_key::CONTENT_LEN)?.u64(1)?;
2731 e.u64(40)?.array(64)?;
2732 for i in 0..64u64 {
2733 e.u64(i)?;
2734 }
2735 Ok(())
2736 });
2737 assert_eq!(DataHeader::decode(&bytes).unwrap().content_len, Some(1));
2738 }
2739
2740 #[test]
2743 fn duplicate_keys_are_rejected() {
2744 let bytes = encode_with(|e| {
2745 e.map(2)?;
2746 e.u64(1)?.u64(1)?;
2747 e.u64(1)?.u64(2)?;
2748 Ok(())
2749 });
2750 assert_eq!(
2751 DataHeader::decode(&bytes).unwrap_err(),
2752 HeaderError::DuplicateKey(1)
2753 );
2754 }
2755
2756 #[test]
2757 fn non_uint_keys_are_rejected() {
2758 let bytes = encode_with(|e| {
2759 e.map(1)?;
2760 e.str("endpoint")?.str("/t")?;
2761 Ok(())
2762 });
2763 assert_eq!(
2764 DataHeader::decode(&bytes).unwrap_err(),
2765 HeaderError::NonUintKey
2766 );
2767
2768 let negative = encode_with(|e| {
2769 e.map(1)?;
2770 e.i64(-1)?.u64(1)?;
2771 Ok(())
2772 });
2773 assert_eq!(
2774 DataHeader::decode(&negative).unwrap_err(),
2775 HeaderError::NonUintKey
2776 );
2777 }
2778
2779 #[test]
2780 fn indefinite_maps_are_rejected() {
2781 let bytes = encode_with(|e| {
2782 e.begin_map()?;
2783 e.u64(1)?.u64(1)?;
2784 e.end()?;
2785 Ok(())
2786 });
2787 assert_eq!(
2788 DataHeader::decode(&bytes).unwrap_err(),
2789 HeaderError::Indefinite
2790 );
2791 }
2792
2793 #[test]
2794 fn indefinite_arrays_are_rejected() {
2795 let bytes = encode_with(|e| {
2796 e.map(5)?;
2797 e.u64(0)?.begin_array()?.u64(0)?.end()?;
2798 e.u64(1)?.u64(1)?;
2799 e.u64(2)?.u64(1)?;
2800 e.u64(3)?.array(0)?;
2801 e.u64(4)?.array(0)?;
2802 Ok(())
2803 });
2804 assert_eq!(Hello::decode(&bytes).unwrap_err(), HeaderError::Indefinite);
2805 }
2806
2807 #[test]
2808 fn value_type_mismatches_are_rejected() {
2809 let bytes = encode_with(|e| {
2810 e.map(1)?;
2811 e.u64(data_key::CONTENT_LEN)?.str("not a number")?;
2812 Ok(())
2813 });
2814 assert!(matches!(
2815 DataHeader::decode(&bytes).unwrap_err(),
2816 HeaderError::Malformed(_)
2817 ));
2818 }
2819
2820 #[test]
2821 fn missing_required_keys_are_rejected() {
2822 let bytes = encode_with(|e| {
2824 e.map(1)?;
2825 e.u64(error_key::MESSAGE)?.str("why")?;
2826 Ok(())
2827 });
2828 assert_eq!(
2829 ErrorHeader::decode(&bytes).unwrap_err(),
2830 HeaderError::MissingKey(error_key::CODE)
2831 );
2832
2833 let bytes = encode_with(|e| {
2835 e.map(4)?;
2836 e.u64(0)?.array(1)?.u64(0)?;
2837 e.u64(1)?.u64(16384)?;
2838 e.u64(2)?.u64(16)?;
2839 e.u64(4)?.array(0)?;
2840 Ok(())
2841 });
2842 assert_eq!(
2843 Hello::decode(&bytes).unwrap_err(),
2844 HeaderError::MissingKey(hello_key::CAPABILITIES)
2845 );
2846 }
2847
2848 #[test]
2851 fn an_empty_filter_is_legal_and_survives_the_roundtrip() {
2852 let h = SubscriptionHeader::new("/md", "");
2853 let bytes = h.encode();
2854 assert_eq!(SubscriptionHeader::decode(&bytes).unwrap(), h);
2855 assert!(
2858 bytes.contains(&0x60),
2859 "the empty filter is encoded: {bytes:?}"
2860 );
2861 }
2862
2863 #[test]
2864 fn the_filter_grammar_accepts_what_docs_protocol_6_4_permits() {
2865 for ok in [
2866 "",
2867 "#",
2868 "px",
2869 "px.eur",
2870 "px.*",
2871 "*.eur",
2872 "sensors.*.temp",
2873 "px.#",
2874 "px.",
2875 "a..b",
2876 ] {
2877 assert_eq!(filter::validate(ok), Ok(()), "{ok:?} must be legal");
2878 }
2879 }
2880
2881 #[test]
2882 fn the_filter_grammar_rejects_partial_and_misplaced_wildcards() {
2883 for bad in [
2884 "px*", "*px", "p*x.eur", "px.e*ur", "#.px", "px.#.eur", "px#",
2885 ] {
2886 assert!(
2887 matches!(filter::validate(bad), Err(HeaderError::InvalidFilter(_))),
2888 "{bad:?} must be rejected"
2889 );
2890 }
2891 }
2892
2893 #[test]
2894 fn an_illegal_filter_is_rejected_at_the_codec_boundary() {
2895 let bytes = SubscriptionHeader::new("/md", "px.#.eur").encode();
2899 assert!(matches!(
2900 SubscriptionHeader::decode(&bytes),
2901 Err(HeaderError::InvalidFilter(_))
2902 ));
2903 let e: Error = HeaderError::InvalidFilter("`#` must be the final segment").into();
2904 assert!(e.to_string().contains("invalid topic filter"));
2905 }
2906
2907 #[test]
2908 fn subscription_strings_are_capped() {
2909 for (key, max) in [
2910 (subscription_key::ENDPOINT, limits::MAX_ENDPOINT_BYTES),
2911 (subscription_key::FILTER, limits::MAX_FILTER_BYTES),
2912 ] {
2913 let build = |len: usize| {
2914 let text = "a".repeat(len);
2915 let mut h = SubscriptionHeader::new("/md", "px.");
2916 if key == subscription_key::ENDPOINT {
2917 h.endpoint = text;
2918 } else {
2919 h.filter = text;
2920 }
2921 h.encode()
2922 };
2923 assert_eq!(
2924 SubscriptionHeader::decode(&build(max + 1)).unwrap_err(),
2925 HeaderError::StringTooLong {
2926 key,
2927 len: max + 1,
2928 max
2929 },
2930 "key {key}"
2931 );
2932 assert!(
2933 SubscriptionHeader::decode(&build(max)).is_ok(),
2934 "key {key} at cap"
2935 );
2936 }
2937 }
2938
2939 #[test]
2940 fn subscription_headers_require_both_keys() {
2941 let only = |key: u64| {
2942 encode_with(|e| {
2943 e.map(1)?;
2944 e.u64(key)?.str("/md")?;
2945 Ok(())
2946 })
2947 };
2948 assert_eq!(
2949 SubscriptionHeader::decode(&only(subscription_key::ENDPOINT)).unwrap_err(),
2950 HeaderError::MissingKey(subscription_key::FILTER)
2951 );
2952 assert_eq!(
2953 SubscriptionHeader::decode(&only(subscription_key::FILTER)).unwrap_err(),
2954 HeaderError::MissingKey(subscription_key::ENDPOINT)
2955 );
2956 }
2957
2958 #[test]
2959 fn subscription_headers_reject_malformed_input() {
2960 assert!(SubscriptionHeader::decode(&[]).is_err());
2961 let mut bytes = SubscriptionHeader::new("/md", "px.").encode();
2963 bytes.push(0xff);
2964 assert_eq!(
2965 SubscriptionHeader::decode(&bytes).unwrap_err(),
2966 HeaderError::TrailingBytes
2967 );
2968 let extended = encode_with(|e| {
2970 e.map(3)?;
2971 e.u64(0)?.str("/md")?;
2972 e.u64(1)?.str("px.")?;
2973 e.u64(40)?.array(2)?.u64(1)?.u64(2)?;
2974 Ok(())
2975 });
2976 assert_eq!(
2977 SubscriptionHeader::decode(&extended).unwrap(),
2978 SubscriptionHeader::new("/md", "px.")
2979 );
2980 }
2981
2982 #[test]
2983 fn oversized_strings_are_rejected_per_field() {
2984 let with_text = |key: u64, text: String| -> Vec<u8> {
2986 let mut h = DataHeader::reply();
2987 match key {
2988 data_key::ENDPOINT => h.endpoint = Some(text),
2989 data_key::CONTENT_TYPE => h.content_type = Some(text),
2990 data_key::TRACEPARENT => h.traceparent = Some(text),
2991 data_key::TRACESTATE => h.tracestate = Some(text),
2992 data_key::TOPIC => h.topic = Some(text),
2993 other => panic!("key {other} is not a text field"),
2994 }
2995 h.encode()
2996 };
2997 let cases: [(u64, usize); 5] = [
2998 (data_key::ENDPOINT, limits::MAX_ENDPOINT_BYTES),
2999 (data_key::CONTENT_TYPE, limits::MAX_CONTENT_TYPE_BYTES),
3000 (data_key::TRACEPARENT, limits::MAX_TRACEPARENT_BYTES),
3001 (data_key::TRACESTATE, limits::MAX_TRACESTATE_BYTES),
3002 (data_key::TOPIC, limits::MAX_TOPIC_BYTES),
3003 ];
3004 for (key, max) in cases {
3005 assert_eq!(
3006 DataHeader::decode(&with_text(key, "a".repeat(max + 1))).unwrap_err(),
3007 HeaderError::StringTooLong {
3008 key,
3009 len: max + 1,
3010 max
3011 },
3012 "key {key}"
3013 );
3014 assert!(
3015 DataHeader::decode(&with_text(key, "a".repeat(max))).is_ok(),
3016 "key {key} at cap"
3017 );
3018 }
3019 }
3020
3021 #[test]
3022 fn unordered_keys_are_rejected() {
3023 let bytes = encode_with(|e| {
3026 e.map(3)?;
3027 e.u64(2)?.str("t")?;
3028 e.u64(1)?.u64(1)?;
3029 e.u64(3)?.str("p")?;
3030 Ok(())
3031 });
3032 assert_eq!(
3033 DataHeader::decode(&bytes).unwrap_err(),
3034 HeaderError::UnorderedKey(1)
3035 );
3036 }
3037
3038 #[test]
3039 fn duplicate_extension_keys_are_rejected() {
3040 let bytes = encode_with(|e| {
3041 e.map(3)?;
3042 e.u64(1)?.u64(1)?;
3043 e.u64(1000)?.u64(1)?;
3044 e.u64(1000)?.u64(2)?;
3045 Ok(())
3046 });
3047 assert_eq!(
3048 DataHeader::decode(&bytes).unwrap_err(),
3049 HeaderError::DuplicateKey(1000)
3050 );
3051 }
3052
3053 #[test]
3054 fn oversized_error_messages_are_rejected() {
3055 let big = "m".repeat(limits::MAX_MESSAGE_BYTES + 1);
3056 let bytes = encode_with(|e| {
3057 e.map(2)?;
3058 e.u64(error_key::CODE)?.u64(2)?;
3059 e.u64(error_key::MESSAGE)?.str(&big)?;
3060 Ok(())
3061 });
3062 assert_eq!(
3063 ErrorHeader::decode(&bytes).unwrap_err(),
3064 HeaderError::StringTooLong {
3065 key: error_key::MESSAGE,
3066 len: limits::MAX_MESSAGE_BYTES + 1,
3067 max: limits::MAX_MESSAGE_BYTES
3068 }
3069 );
3070 }
3071
3072 #[test]
3073 fn oversized_lists_are_rejected_without_allocating() {
3074 let bytes = encode_with(|e| {
3076 e.map(1)?;
3077 e.u64(0)?.array(u64::from(u32::MAX))?;
3078 Ok(())
3079 });
3080 assert_eq!(
3081 Hello::decode(&bytes).unwrap_err(),
3082 HeaderError::ListTooLong {
3083 key: hello_key::VERSIONS,
3084 len: u64::from(u32::MAX),
3085 max: limits::MAX_LIST_ITEMS
3086 }
3087 );
3088 }
3089
3090 #[test]
3091 fn lists_exactly_at_the_cap_are_accepted() {
3092 let bytes = encode_with(|e| {
3093 e.map(5)?;
3094 e.u64(0)?.array(limits::MAX_LIST_ITEMS as u64)?;
3095 for i in 0..limits::MAX_LIST_ITEMS as u64 {
3096 e.u64(i)?;
3097 }
3098 e.u64(1)?.u64(16384)?;
3099 e.u64(2)?.u64(16)?;
3100 e.u64(3)?.array(0)?;
3101 e.u64(4)?.array(0)?;
3102 Ok(())
3103 });
3104 assert_eq!(
3105 Hello::decode(&bytes).unwrap().versions.len(),
3106 limits::MAX_LIST_ITEMS
3107 );
3108 }
3109
3110 #[test]
3111 fn trailing_bytes_are_rejected() {
3112 let mut bytes = ErrorHeader::new(ErrorCode::Rejected).encode();
3113 bytes.push(0xff);
3114 assert_eq!(
3115 ErrorHeader::decode(&bytes).unwrap_err(),
3116 HeaderError::TrailingBytes
3117 );
3118 }
3119
3120 #[test]
3121 fn truncated_headers_are_rejected() {
3122 let full = DataHeader::addressed("/t").encode();
3123 for cut in 0..full.len() {
3124 assert!(
3125 DataHeader::decode(&full[..cut]).is_err(),
3126 "prefix of {cut} bytes must not decode"
3127 );
3128 }
3129 }
3130
3131 #[test]
3132 fn empty_input_is_rejected_for_every_header() {
3133 assert!(Hello::decode(&[]).is_err());
3134 assert!(DataHeader::decode(&[]).is_err());
3135 assert!(ErrorHeader::decode(&[]).is_err());
3136 assert!(SubscriptionHeader::decode(&[]).is_err());
3137 }
3138
3139 #[test]
3140 fn tags_are_rejected() {
3141 let bytes = encode_with(|e| {
3144 e.map(1)?;
3145 e.u64(50)?.tag(minicbor::data::IanaTag::Cbor)?.u64(1)?;
3146 Ok(())
3147 });
3148 assert_eq!(
3149 DataHeader::decode(&bytes).unwrap_err(),
3150 HeaderError::Malformed("tags are not allowed")
3151 );
3152 }
3153
3154 #[test]
3157 fn unknown_error_codes_survive_decoding() {
3158 let err = ErrorHeader {
3159 code: 99,
3160 message: None,
3161 };
3162 let decoded = ErrorHeader::decode(&err.encode()).unwrap();
3163 assert_eq!(decoded.code, 99);
3164 assert_eq!(decoded.error_code(), None);
3165 }
3166
3167 #[test]
3168 fn header_errors_become_protocol_errors() {
3169 let e: Error = HeaderError::DuplicateKey(3).into();
3170 assert!(matches!(e, Error::Protocol(_)));
3171 assert!(e.to_string().contains("duplicate header key 3"));
3172 }
3173}