Skip to main content

s2_common/record/
mod.rs

1mod command;
2mod envelope;
3mod fencing;
4mod metering;
5
6use bytes::Bytes;
7pub use command::{CommandOp, CommandPayloadError, CommandRecord};
8pub use envelope::{EnvelopeRecord, HeaderValidationError};
9pub use fencing::{FencingToken, FencingTokenTooLongError, MAX_FENCING_TOKEN_LENGTH};
10pub use metering::{Metered, MeteredExt, MeteredSize};
11
12use crate::deep_size::DeepSize;
13
14pub type SeqNum = u64;
15pub type NonZeroSeqNum = std::num::NonZeroU64;
16pub type Timestamp = u64;
17
18#[derive(Debug, PartialEq, Eq, Clone, Copy)]
19pub struct StreamPosition {
20    pub seq_num: SeqNum,
21    pub timestamp: Timestamp,
22}
23
24impl StreamPosition {
25    pub const MIN: StreamPosition = StreamPosition {
26        seq_num: SeqNum::MIN,
27        timestamp: Timestamp::MIN,
28    };
29}
30
31impl std::fmt::Display for StreamPosition {
32    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
33        write!(f, "{} @ {}", self.seq_num, self.timestamp)
34    }
35}
36
37impl DeepSize for StreamPosition {
38    fn deep_size(&self) -> usize {
39        self.seq_num.deep_size() + self.timestamp.deep_size()
40    }
41}
42
43#[derive(Debug, PartialEq, thiserror::Error)]
44pub enum RecordPartsError {
45    #[error("unknown command")]
46    UnknownCommand,
47    #[error("invalid `{0}` command: {1}")]
48    CommandPayload(CommandOp, CommandPayloadError),
49    #[error("invalid header: {0}")]
50    Header(#[from] HeaderValidationError),
51}
52
53#[derive(Debug, Clone, PartialEq, Eq)]
54pub struct Header {
55    pub name: Bytes,
56    pub value: Bytes,
57}
58
59impl DeepSize for Header {
60    fn deep_size(&self) -> usize {
61        self.name.len() + self.value.len()
62    }
63}
64
65impl MeteredSize for Record {
66    fn metered_size(&self) -> usize {
67        match self {
68            Self::Command(command) => command.metered_size(),
69            Self::Envelope(envelope) => envelope.metered_size(),
70        }
71    }
72}
73
74#[derive(Debug, PartialEq, Eq, Clone)]
75pub enum Record {
76    Command(CommandRecord),
77    Envelope(EnvelopeRecord),
78}
79
80impl DeepSize for Record {
81    fn deep_size(&self) -> usize {
82        match self {
83            Self::Command(c) => c.deep_size(),
84            Self::Envelope(e) => e.deep_size(),
85        }
86    }
87}
88
89impl Record {
90    pub fn try_from_parts(headers: Vec<Header>, body: Bytes) -> Result<Self, RecordPartsError> {
91        if headers.len() == 1 {
92            let header = &headers[0];
93            if header.name.is_empty() {
94                let op = CommandOp::from_id(header.value.as_ref())
95                    .ok_or(RecordPartsError::UnknownCommand)?;
96                let command_record = CommandRecord::try_from_parts(op, body.as_ref())
97                    .map_err(|e| RecordPartsError::CommandPayload(op, e))?;
98                return Ok(Self::Command(command_record));
99            }
100        }
101        let envelope = EnvelopeRecord::try_from_parts(headers, body)?;
102        Ok(Self::Envelope(envelope))
103    }
104
105    pub fn into_parts(self) -> (Vec<Header>, Bytes) {
106        match self {
107            Record::Envelope(e) => e.into_parts(),
108            Record::Command(c) => {
109                let op = c.op();
110                let header = Header {
111                    name: Bytes::new(),
112                    value: Bytes::from_static(op.to_id()),
113                };
114                (vec![header], c.payload())
115            }
116        }
117    }
118}
119
120#[derive(Debug, Clone, PartialEq, Eq)]
121pub struct Sequenced<T> {
122    position: StreamPosition,
123    inner: T,
124}
125
126impl<T> Sequenced<T> {
127    pub const fn new(position: StreamPosition, inner: T) -> Self {
128        Self { position, inner }
129    }
130
131    pub const fn position(&self) -> &StreamPosition {
132        &self.position
133    }
134
135    pub fn inner(&self) -> &T {
136        &self.inner
137    }
138
139    pub fn as_ref(&self) -> Sequenced<&T> {
140        Sequenced::new(self.position, &self.inner)
141    }
142
143    pub fn parts(&self) -> (StreamPosition, &T) {
144        (self.position, &self.inner)
145    }
146
147    pub fn into_parts(self) -> (StreamPosition, T) {
148        (self.position, self.inner)
149    }
150}
151
152pub type SequencedRecord = Sequenced<Record>;
153
154impl<T> MeteredSize for Sequenced<T>
155where
156    T: MeteredSize,
157{
158    fn metered_size(&self) -> usize {
159        self.inner.metered_size()
160    }
161}
162
163impl<T> DeepSize for Sequenced<T>
164where
165    T: DeepSize,
166{
167    fn deep_size(&self) -> usize {
168        self.position.deep_size() + self.inner.deep_size()
169    }
170}
171
172impl<T> Metered<T>
173where
174    T: MeteredSize,
175{
176    pub fn sequenced(self, position: StreamPosition) -> Metered<Sequenced<T>> {
177        Metered::with_size(
178            self.metered_size(),
179            Sequenced::new(position, self.into_inner()),
180        )
181    }
182}
183
184impl<T> Metered<Sequenced<T>> {
185    pub fn parts(&self) -> (StreamPosition, Metered<&T>) {
186        let size = self.metered_size();
187        let (position, inner) = self.as_ref().into_inner().parts();
188        (position, Metered::with_size(size, inner))
189    }
190
191    pub fn into_parts(self) -> (StreamPosition, Metered<T>) {
192        let size = self.metered_size();
193        let (position, inner) = self.into_inner().into_parts();
194        (position, Metered::with_size(size, inner))
195    }
196}
197
198#[cfg(test)]
199mod test {
200    use rstest::rstest;
201
202    use super::*;
203
204    fn semantic_metered_size(record: &Record) -> usize {
205        let (headers, body) = record.clone().into_parts();
206        8 + (2 * headers.len())
207            + headers
208                .iter()
209                .map(|header| header.name.len() + header.value.len())
210                .sum::<usize>()
211            + body.len()
212    }
213
214    #[test]
215    fn empty_header_name_solo() {
216        let headers = vec![Header {
217            name: Bytes::new(),
218            value: Bytes::from("hi"),
219        }];
220        let body = Bytes::from("hello");
221        assert_eq!(
222            Record::try_from_parts(headers, body),
223            Err(RecordPartsError::UnknownCommand)
224        );
225    }
226
227    #[test]
228    fn empty_header_name_among_others() {
229        let headers = vec![
230            Header {
231                name: Bytes::from("boku"),
232                value: Bytes::from("hi"),
233            },
234            Header {
235                name: Bytes::new(),
236                value: Bytes::from("hi"),
237            },
238        ];
239        let body = Bytes::from("hello");
240        assert_eq!(
241            Record::try_from_parts(headers, body),
242            Err(RecordPartsError::Header(HeaderValidationError::NameEmpty))
243        );
244    }
245
246    fn command_parts(op: &'static [u8], payload: &'static [u8]) -> (Vec<Header>, Bytes) {
247        let headers = vec![Header {
248            name: Bytes::new(),
249            value: Bytes::from_static(op),
250        }];
251        let body = Bytes::from_static(payload);
252        (headers, body)
253    }
254
255    fn assert_valid_command_record(op: &'static [u8], payload: &'static [u8]) {
256        let (headers, body) = command_parts(op, payload);
257        let record = Record::try_from_parts(headers.clone(), body.clone()).unwrap();
258        let record_metered = record.metered_size();
259        match &record {
260            Record::Command(cmd) => {
261                assert_eq!(cmd.op().to_id(), op);
262                assert_eq!(cmd.payload().as_ref(), payload);
263            }
264            other => panic!("Command expected, got {other:?}"),
265        }
266        assert_eq!(record_metered, semantic_metered_size(&record));
267    }
268
269    #[rstest]
270    #[case::fence_empty(b"fence", b"")]
271    #[case::fence_uuid(b"fence", b"my-special-uuid")]
272    #[case::trim_0(b"trim", b"\x00\x00\x00\x00\x00\x00\x00\x00")]
273    fn valid_command_records(#[case] op: &'static [u8], #[case] payload: &'static [u8]) {
274        assert_valid_command_record(op, payload);
275    }
276
277    #[rstest]
278    #[case::fence_too_long(
279        b"fence",
280        b"toolongtoolongtoolongtoolongtoolongtoolongtoolong",
281        RecordPartsError::CommandPayload(
282            CommandOp::Fence,
283            CommandPayloadError::FencingTokenTooLong(FencingTokenTooLongError(49)),
284        )
285    )]
286    #[case::trim_empty(
287        b"trim",
288        b"",
289        RecordPartsError::CommandPayload(CommandOp::Trim, CommandPayloadError::TrimPointSize(0),)
290    )]
291    #[case::trim_overflow(
292        b"trim",
293        b"\x00\x00\x00\x00\x00\x00\x00\x00\x00",
294        RecordPartsError::CommandPayload(CommandOp::Trim, CommandPayloadError::TrimPointSize(9),)
295    )]
296    fn invalid_command_records(
297        #[case] op: &'static [u8],
298        #[case] payload: &'static [u8],
299        #[case] expected: RecordPartsError,
300    ) {
301        let (headers, body) = command_parts(op, payload);
302        assert_eq!(Record::try_from_parts(headers, body), Err(expected));
303    }
304}