Skip to main content

s2_common/
stream.rs

1use std::{marker::PhantomData, ops::Deref, str::FromStr, time::Duration};
2
3use compact_str::{CompactString, ToCompactString};
4use time::OffsetDateTime;
5
6use super::{
7    ValidationError,
8    strings::{NameProps, PrefixProps, StartAfterProps, StrProps},
9};
10use crate::{
11    caps,
12    encryption::EncryptionAlgorithm,
13    read_extent::{ReadLimit, ReadUntil},
14    record::{
15        FencingToken, Metered, MeteredSize, Record, SeqNum, Sequenced, StreamPosition, Timestamp,
16    },
17    resources::ListItemsRequest,
18};
19
20#[derive(Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
21#[cfg_attr(
22    feature = "rkyv",
23    derive(rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)
24)]
25pub struct StreamNameStr<T: StrProps>(CompactString, PhantomData<T>);
26
27impl<T: StrProps> StreamNameStr<T> {
28    fn validate_str(name: &str) -> Result<(), ValidationError> {
29        if !T::IS_PREFIX && name.is_empty() {
30            return Err(format!("stream {} must not be empty", T::FIELD_NAME).into());
31        }
32
33        if !T::IS_PREFIX && (name == "." || name == "..") {
34            return Err(format!("stream {} must not be \".\" or \"..\"", T::FIELD_NAME).into());
35        }
36
37        if name.contains('\0') {
38            return Err(format!("stream {} must not contain NUL bytes", T::FIELD_NAME).into());
39        }
40
41        if name.len() > caps::MAX_STREAM_NAME_LEN {
42            return Err(format!(
43                "stream {} must not exceed {} bytes in length",
44                T::FIELD_NAME,
45                caps::MAX_STREAM_NAME_LEN
46            )
47            .into());
48        }
49
50        Ok(())
51    }
52}
53
54#[cfg(feature = "utoipa")]
55impl<T> utoipa::PartialSchema for StreamNameStr<T>
56where
57    T: StrProps,
58{
59    fn schema() -> utoipa::openapi::RefOr<utoipa::openapi::schema::Schema> {
60        utoipa::openapi::Object::builder()
61            .schema_type(utoipa::openapi::Type::String)
62            .min_length((!T::IS_PREFIX).then_some(caps::MIN_STREAM_NAME_LEN))
63            .max_length(Some(caps::MAX_STREAM_NAME_LEN))
64            .into()
65    }
66}
67
68#[cfg(feature = "utoipa")]
69impl<T> utoipa::ToSchema for StreamNameStr<T> where T: StrProps {}
70
71impl<T: StrProps> serde::Serialize for StreamNameStr<T> {
72    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
73    where
74        S: serde::Serializer,
75    {
76        serializer.serialize_str(&self.0)
77    }
78}
79
80impl<'de, T: StrProps> serde::Deserialize<'de> for StreamNameStr<T> {
81    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
82    where
83        D: serde::Deserializer<'de>,
84    {
85        let s = CompactString::deserialize(deserializer)?;
86        s.try_into().map_err(serde::de::Error::custom)
87    }
88}
89
90impl<T: StrProps> AsRef<str> for StreamNameStr<T> {
91    fn as_ref(&self) -> &str {
92        &self.0
93    }
94}
95
96impl<T: StrProps> Deref for StreamNameStr<T> {
97    type Target = str;
98
99    fn deref(&self) -> &Self::Target {
100        &self.0
101    }
102}
103
104impl<T: StrProps> TryFrom<CompactString> for StreamNameStr<T> {
105    type Error = ValidationError;
106
107    fn try_from(name: CompactString) -> Result<Self, Self::Error> {
108        Self::validate_str(&name)?;
109        Ok(Self(name, PhantomData))
110    }
111}
112
113impl<T: StrProps> FromStr for StreamNameStr<T> {
114    type Err = ValidationError;
115
116    fn from_str(s: &str) -> Result<Self, Self::Err> {
117        Self::validate_str(s)?;
118        Ok(Self(s.to_compact_string(), PhantomData))
119    }
120}
121
122impl<T: StrProps> std::fmt::Debug for StreamNameStr<T> {
123    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
124        f.write_str(&self.0)
125    }
126}
127
128impl<T: StrProps> std::fmt::Display for StreamNameStr<T> {
129    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
130        f.write_str(&self.0)
131    }
132}
133
134impl<T: StrProps> From<StreamNameStr<T>> for CompactString {
135    fn from(value: StreamNameStr<T>) -> Self {
136        value.0
137    }
138}
139
140pub type StreamName = StreamNameStr<NameProps>;
141
142pub type StreamNamePrefix = StreamNameStr<PrefixProps>;
143
144impl Default for StreamNamePrefix {
145    fn default() -> Self {
146        StreamNameStr(CompactString::default(), PhantomData)
147    }
148}
149
150impl From<StreamName> for StreamNamePrefix {
151    fn from(value: StreamName) -> Self {
152        Self(value.0, PhantomData)
153    }
154}
155
156pub type StreamNameStartAfter = StreamNameStr<StartAfterProps>;
157
158impl Default for StreamNameStartAfter {
159    fn default() -> Self {
160        StreamNameStr(CompactString::default(), PhantomData)
161    }
162}
163
164impl From<StreamName> for StreamNameStartAfter {
165    fn from(value: StreamName) -> Self {
166        Self(value.0, PhantomData)
167    }
168}
169
170#[derive(Debug, Clone)]
171pub struct StreamInfo {
172    pub name: StreamName,
173    pub created_at: OffsetDateTime,
174    pub deleted_at: Option<OffsetDateTime>,
175    pub cipher: Option<EncryptionAlgorithm>,
176}
177
178#[derive(Debug, Clone)]
179pub struct AppendRecord<T = Record>(AppendRecordParts<T>);
180
181impl<T> AppendRecord<T> {
182    pub fn parts(&self) -> &AppendRecordParts<T> {
183        let Self(parts) = self;
184        parts
185    }
186
187    pub fn into_parts(self) -> AppendRecordParts<T> {
188        let Self(parts) = self;
189        parts
190    }
191}
192
193impl<T> MeteredSize for AppendRecord<T> {
194    fn metered_size(&self) -> usize {
195        self.0.record.metered_size()
196    }
197}
198
199#[derive(Debug, Clone)]
200pub struct AppendRecordParts<T = Record> {
201    pub timestamp: Option<Timestamp>,
202    pub record: Metered<T>,
203}
204
205impl<T> MeteredSize for AppendRecordParts<T> {
206    fn metered_size(&self) -> usize {
207        self.record.metered_size()
208    }
209}
210
211impl<T> From<AppendRecord<T>> for AppendRecordParts<T> {
212    fn from(record: AppendRecord<T>) -> Self {
213        record.into_parts()
214    }
215}
216
217impl<T> TryFrom<AppendRecordParts<T>> for AppendRecord<T> {
218    type Error = &'static str;
219
220    fn try_from(parts: AppendRecordParts<T>) -> Result<Self, Self::Error> {
221        if parts.metered_size() > caps::RECORD_BATCH_MAX.bytes {
222            Err("record must have metered size less than 1 MiB")
223        } else {
224            Ok(Self(parts))
225        }
226    }
227}
228
229#[derive(Clone)]
230pub struct AppendRecordBatch<T = Record>(Metered<Vec<AppendRecord<T>>>);
231
232impl<T> std::fmt::Debug for AppendRecordBatch<T> {
233    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
234        f.debug_struct("AppendRecordBatch")
235            .field("num_records", &self.0.len())
236            .field("metered_size", &self.0.metered_size())
237            .finish()
238    }
239}
240
241impl<T> MeteredSize for AppendRecordBatch<T> {
242    fn metered_size(&self) -> usize {
243        self.0.metered_size()
244    }
245}
246
247impl<T> std::ops::Deref for AppendRecordBatch<T> {
248    type Target = [AppendRecord<T>];
249
250    fn deref(&self) -> &Self::Target {
251        &self.0
252    }
253}
254
255impl<T> TryFrom<Metered<Vec<AppendRecord<T>>>> for AppendRecordBatch<T> {
256    type Error = &'static str;
257
258    fn try_from(records: Metered<Vec<AppendRecord<T>>>) -> Result<Self, Self::Error> {
259        if records.is_empty() {
260            return Err("record batch must not be empty");
261        }
262
263        if records.len() > caps::RECORD_BATCH_MAX.count {
264            return Err("record batch must not exceed 1000 records");
265        }
266
267        if records.metered_size() > caps::RECORD_BATCH_MAX.bytes {
268            return Err("record batch must not exceed a metered size of 1 MiB");
269        }
270
271        Ok(Self(records))
272    }
273}
274
275impl<T> TryFrom<Vec<AppendRecord<T>>> for AppendRecordBatch<T> {
276    type Error = &'static str;
277
278    fn try_from(records: Vec<AppendRecord<T>>) -> Result<Self, Self::Error> {
279        let records = Metered::from(records);
280        Self::try_from(records)
281    }
282}
283
284impl<T> IntoIterator for AppendRecordBatch<T> {
285    type Item = AppendRecord<T>;
286    type IntoIter = std::vec::IntoIter<Self::Item>;
287
288    fn into_iter(self) -> Self::IntoIter {
289        self.0.into_iter()
290    }
291}
292
293#[derive(Debug, Clone)]
294pub struct AppendInput<T = Record> {
295    pub records: AppendRecordBatch<T>,
296    pub match_seq_num: Option<SeqNum>,
297    pub fencing_token: Option<FencingToken>,
298}
299
300#[derive(Debug, Clone)]
301pub struct AppendAck {
302    pub start: StreamPosition,
303    pub end: StreamPosition,
304    pub tail: StreamPosition,
305}
306
307#[derive(Debug, Clone, Copy, PartialEq, Eq)]
308pub enum ReadPosition {
309    SeqNum(SeqNum),
310    Timestamp(Timestamp),
311}
312
313#[derive(Debug, Clone, Copy)]
314pub enum ReadFrom {
315    SeqNum(SeqNum),
316    Timestamp(Timestamp),
317    TailOffset(u64),
318}
319
320impl Default for ReadFrom {
321    fn default() -> Self {
322        Self::SeqNum(0)
323    }
324}
325
326#[derive(Debug, Default, Clone, Copy)]
327pub struct ReadStart {
328    pub from: ReadFrom,
329    pub clamp: bool,
330}
331
332#[derive(Debug, Default, Clone, Copy)]
333pub struct ReadEnd {
334    pub limit: ReadLimit,
335    pub until: ReadUntil,
336    pub wait: Option<Duration>,
337}
338
339impl ReadEnd {
340    pub fn may_follow(&self) -> bool {
341        (self.limit.is_unbounded() && self.until.is_unbounded())
342            || self.wait.is_some_and(|d| d > Duration::ZERO)
343    }
344}
345
346#[derive(Clone)]
347pub struct ReadBatch<T = Record> {
348    pub records: Metered<Vec<Sequenced<T>>>,
349    pub tail: Option<StreamPosition>,
350}
351
352impl<T> Default for ReadBatch<T>
353where
354    T: MeteredSize,
355{
356    fn default() -> Self {
357        Self {
358            records: Metered::default(),
359            tail: None,
360        }
361    }
362}
363
364impl<T> std::fmt::Debug for ReadBatch<T> {
365    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
366        f.debug_struct("ReadBatch")
367            .field("num_records", &self.records.len())
368            .field("metered_size", &self.records.metered_size())
369            .field("tail", &self.tail)
370            .finish()
371    }
372}
373
374#[derive(Debug, Clone)]
375pub enum ReadSessionOutput<T = Record> {
376    Heartbeat(StreamPosition),
377    Batch(ReadBatch<T>),
378}
379
380pub type ListStreamsRequest = ListItemsRequest<StreamNamePrefix, StreamNameStartAfter>;
381
382#[cfg(test)]
383mod test {
384    use rstest::rstest;
385
386    use super::{
387        super::strings::{NameProps, PrefixProps, StartAfterProps},
388        *,
389    };
390
391    #[rstest]
392    #[case::normal("my-stream".to_owned())]
393    #[case::control_chars("a\tb\nc\rd\x01e".to_owned())]
394    #[case::unicode("stream/名前 😀?#%20".to_owned())]
395    #[case::max_len("a".repeat(crate::caps::MAX_STREAM_NAME_LEN))]
396    fn validate_name_ok(#[case] name: String) {
397        assert_eq!(StreamNameStr::<NameProps>::validate_str(&name), Ok(()));
398    }
399
400    #[rstest]
401    #[case::empty("".to_owned())]
402    #[case::dot(".".to_owned())]
403    #[case::dot_dot("..".to_owned())]
404    #[case::too_long("a".repeat(crate::caps::MAX_STREAM_NAME_LEN + 1))]
405    #[case::nul("a\0b".to_owned())]
406    #[case::leading_nul("\0a".to_owned())]
407    #[case::trailing_nul("a\0".to_owned())]
408    #[case::only_nul("\0".to_owned())]
409    fn validate_name_err(#[case] name: String) {
410        StreamNameStr::<NameProps>::validate_str(&name).expect_err("expected validation error");
411    }
412
413    #[rstest]
414    #[case::empty("".to_owned())]
415    #[case::dot(".".to_owned())]
416    #[case::dot_dot("..".to_owned())]
417    #[case::max_len("a".repeat(crate::caps::MAX_STREAM_NAME_LEN))]
418    fn validate_prefix_ok(#[case] prefix: String) {
419        assert_eq!(StreamNameStr::<PrefixProps>::validate_str(&prefix), Ok(()));
420    }
421
422    #[rstest]
423    #[case::too_long("a".repeat(crate::caps::MAX_STREAM_NAME_LEN + 1))]
424    #[case::nul("a\0b".to_owned())]
425    #[case::only_nul("\0".to_owned())]
426    fn validate_prefix_err(#[case] prefix: String) {
427        StreamNameStr::<PrefixProps>::validate_str(&prefix).expect_err("expected validation error");
428    }
429
430    #[rstest]
431    #[case::empty("".to_owned())]
432    #[case::dot(".".to_owned())]
433    #[case::dot_dot("..".to_owned())]
434    #[case::max_len("a".repeat(crate::caps::MAX_STREAM_NAME_LEN))]
435    fn validate_start_after_ok(#[case] start_after: String) {
436        assert_eq!(
437            StreamNameStr::<StartAfterProps>::validate_str(&start_after),
438            Ok(())
439        );
440    }
441
442    #[rstest]
443    #[case::too_long("a".repeat(crate::caps::MAX_STREAM_NAME_LEN + 1))]
444    #[case::nul("a\0b".to_owned())]
445    #[case::only_nul("\0".to_owned())]
446    fn validate_start_after_err(#[case] start_after: String) {
447        StreamNameStr::<StartAfterProps>::validate_str(&start_after)
448            .expect_err("expected validation error");
449    }
450
451    #[test]
452    fn append_record_batch_rejects_empty_batches() {
453        let empty_batch: Result<AppendRecordBatch, _> = Vec::<AppendRecord>::new().try_into();
454
455        assert_eq!(empty_batch.unwrap_err(), "record batch must not be empty");
456    }
457}