Skip to main content

ntex_mqtt/
topic.rs

1use std::{fmt, fmt::Write, io};
2
3use ntex_bytes::ByteString;
4
5#[allow(clippy::match_same_arms)]
6pub(crate) fn is_valid(topic: &str) -> bool {
7    if topic.is_empty() {
8        false
9    } else {
10        enum PrevState {
11            None,
12            LevelSep,
13            SingleWildcard,
14            MultiWildcard,
15            Other,
16        }
17
18        let mut previous = PrevState::None;
19        for current in topic.bytes() {
20            previous = match (current, &previous) {
21                (_, PrevState::MultiWildcard) => return false, // `#` is not last char
22                (b'+', PrevState::None | PrevState::LevelSep) => PrevState::SingleWildcard,
23                (b'#', PrevState::None | PrevState::LevelSep) => PrevState::MultiWildcard,
24                (b'+' | b'#', _) => return false, // `+` or `#` after char other than `/`
25                (b'/', _) => PrevState::LevelSep,
26                (_, PrevState::SingleWildcard) => return false, // `+` is followed by char other than `/`
27                _ => PrevState::Other,
28            }
29        }
30        true
31    }
32}
33
34#[derive(Copy, Clone, Debug, PartialEq, Eq)]
35pub enum TopicFilterError {
36    InvalidTopic,
37    InvalidLevel,
38}
39
40#[derive(Debug, Clone, Hash, Eq, PartialEq, serde::Serialize, serde::Deserialize)]
41pub enum TopicFilterLevel {
42    Normal(ByteString),
43    System(ByteString),
44    Blank,
45    SingleWildcard, // Single level wildcard +
46    MultiWildcard,  // Multi-level wildcard #
47}
48
49impl TopicFilterLevel {
50    fn is_valid(&self) -> bool {
51        match *self {
52            TopicFilterLevel::Normal(ref s) | TopicFilterLevel::System(ref s) => {
53                !s.contains(['+', '#'])
54            }
55            _ => true,
56        }
57    }
58}
59
60fn match_topic<T: MatchLevel, L: Iterator<Item = T>>(
61    superset: &TopicFilter,
62    subset: L,
63) -> bool {
64    let mut superset = superset.0.iter();
65
66    for (index, subset_level) in subset.enumerate() {
67        match superset.next() {
68            Some(TopicFilterLevel::SingleWildcard) => {
69                if !subset_level.match_level(&TopicFilterLevel::SingleWildcard, index) {
70                    return false;
71                }
72            }
73            Some(TopicFilterLevel::MultiWildcard) => {
74                return subset_level.match_level(&TopicFilterLevel::MultiWildcard, index);
75            }
76            Some(level) if subset_level.match_level(level, index) => (),
77            _ => return false,
78        }
79    }
80
81    match superset.next() {
82        Some(&TopicFilterLevel::MultiWildcard) | None => true,
83        Some(_) => false,
84    }
85}
86
87#[derive(Debug, Clone, Hash, Eq, PartialEq, serde::Serialize, serde::Deserialize)]
88pub struct TopicFilter(Vec<TopicFilterLevel>);
89
90impl TopicFilter {
91    pub fn levels(&self) -> &[TopicFilterLevel] {
92        &self.0
93    }
94
95    fn is_valid(&self) -> bool {
96        self.0
97            .iter()
98            .position(|level| !level.is_valid())
99            .or_else(|| {
100                self.0.iter().enumerate().position(|(pos, level)| match *level {
101                    TopicFilterLevel::MultiWildcard => pos != self.0.len() - 1,
102                    TopicFilterLevel::System(_) => pos != 0,
103                    _ => false,
104                })
105            })
106            .is_none()
107    }
108
109    pub fn matches_filter(&self, topic: &TopicFilter) -> bool {
110        match_topic(self, topic.0.iter())
111    }
112
113    pub fn matches_topic<S: AsRef<str> + ?Sized>(&self, topic: &S) -> bool {
114        match_topic(self, topic.as_ref().split('/'))
115    }
116}
117
118impl TryFrom<&[TopicFilterLevel]> for TopicFilter {
119    type Error = TopicFilterError;
120
121    fn try_from(s: &[TopicFilterLevel]) -> Result<Self, Self::Error> {
122        let mut v = vec![];
123        v.extend_from_slice(s);
124
125        TopicFilter::try_from(v)
126    }
127}
128
129impl TryFrom<Vec<TopicFilterLevel>> for TopicFilter {
130    type Error = TopicFilterError;
131
132    fn try_from(v: Vec<TopicFilterLevel>) -> Result<Self, Self::Error> {
133        let tf = TopicFilter(v);
134        if tf.is_valid() {
135            Ok(tf)
136        } else {
137            Err(TopicFilterError::InvalidTopic)
138        }
139    }
140}
141
142impl From<TopicFilter> for Vec<TopicFilterLevel> {
143    fn from(t: TopicFilter) -> Self {
144        t.0
145    }
146}
147
148trait MatchLevel {
149    fn match_level(&self, level: &TopicFilterLevel, index: usize) -> bool;
150}
151
152impl MatchLevel for TopicFilterLevel {
153    fn match_level(&self, level: &TopicFilterLevel, index: usize) -> bool {
154        match_level_impl(self, level, index)
155    }
156}
157
158impl MatchLevel for &TopicFilterLevel {
159    fn match_level(&self, level: &TopicFilterLevel, index: usize) -> bool {
160        match_level_impl(self, level, index)
161    }
162}
163
164fn match_level_impl(
165    subset_level: &TopicFilterLevel,
166    superset_level: &TopicFilterLevel,
167    _index: usize,
168) -> bool {
169    match superset_level {
170        TopicFilterLevel::Normal(rhs) => {
171            matches!(subset_level, TopicFilterLevel::Normal(lhs) if lhs == rhs)
172        }
173        TopicFilterLevel::System(rhs) => {
174            matches!(subset_level, TopicFilterLevel::System(lhs) if lhs == rhs)
175        }
176        TopicFilterLevel::Blank => *subset_level == TopicFilterLevel::Blank,
177        TopicFilterLevel::SingleWildcard => *subset_level != TopicFilterLevel::MultiWildcard,
178        TopicFilterLevel::MultiWildcard => true,
179    }
180}
181
182impl<T: AsRef<str>> MatchLevel for T {
183    fn match_level(&self, level: &TopicFilterLevel, index: usize) -> bool {
184        match level {
185            TopicFilterLevel::Normal(lhs) => lhs == self.as_ref(),
186            TopicFilterLevel::System(lhs) => is_system(self) && lhs == self.as_ref(),
187            TopicFilterLevel::Blank => self.as_ref().is_empty(),
188            TopicFilterLevel::SingleWildcard | TopicFilterLevel::MultiWildcard => {
189                !(index == 0 && is_system(self))
190            }
191        }
192    }
193}
194
195impl TryFrom<ByteString> for TopicFilter {
196    type Error = TopicFilterError;
197
198    fn try_from(value: ByteString) -> Result<Self, Self::Error> {
199        if value.is_empty() {
200            return Err(TopicFilterError::InvalidTopic);
201        }
202
203        value
204            .split('/')
205            .enumerate()
206            .map(|(idx, level)| match level {
207                "+" => Ok(TopicFilterLevel::SingleWildcard),
208                "#" => Ok(TopicFilterLevel::MultiWildcard),
209                "" => Ok(TopicFilterLevel::Blank),
210                _ => {
211                    if level.contains(['+', '#']) {
212                        Err(TopicFilterError::InvalidLevel)
213                    } else if idx == 0 && is_system(level) {
214                        Ok(TopicFilterLevel::System(recover_bstr(&value, level)))
215                    } else {
216                        Ok(TopicFilterLevel::Normal(recover_bstr(&value, level)))
217                    }
218                }
219            })
220            .collect::<Result<Vec<_>, TopicFilterError>>()
221            .map(TopicFilter)
222            .and_then(|topic| {
223                if topic.is_valid() {
224                    Ok(topic)
225                } else {
226                    Err(TopicFilterError::InvalidTopic)
227                }
228            })
229    }
230}
231
232impl std::str::FromStr for TopicFilter {
233    type Err = TopicFilterError;
234
235    fn from_str(value: &str) -> Result<Self, Self::Err> {
236        let s: ByteString = value.into();
237        TopicFilter::try_from(s)
238    }
239}
240
241impl fmt::Display for TopicFilterLevel {
242    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
243        match self {
244            TopicFilterLevel::Normal(s) | TopicFilterLevel::System(s) => {
245                f.write_str(s.as_str())
246            }
247            TopicFilterLevel::Blank => Ok(()),
248            TopicFilterLevel::SingleWildcard => f.write_char('+'),
249            TopicFilterLevel::MultiWildcard => f.write_char('#'),
250        }
251    }
252}
253
254impl fmt::Display for TopicFilter {
255    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
256        let mut iter = self.0.iter();
257        let mut level = iter.next().unwrap();
258        loop {
259            level.fmt(f)?;
260            if let Some(l) = iter.next() {
261                level = l;
262                f.write_char('/')?;
263            } else {
264                break;
265            }
266        }
267        Ok(())
268    }
269}
270
271#[allow(dead_code)]
272pub(crate) trait WriteTopicExt: io::Write {
273    fn write_level(&mut self, level: &TopicFilterLevel) -> io::Result<usize> {
274        match *level {
275            TopicFilterLevel::Normal(ref s) | TopicFilterLevel::System(ref s) => {
276                self.write(s.as_str().as_bytes())
277            }
278            TopicFilterLevel::Blank => Ok(0),
279            TopicFilterLevel::SingleWildcard => self.write(b"+"),
280            TopicFilterLevel::MultiWildcard => self.write(b"#"),
281        }
282    }
283
284    fn write_topic(&mut self, topic: &TopicFilter) -> io::Result<usize> {
285        let mut n = 0;
286        let mut iter = topic.0.iter();
287        let mut level = iter.next().unwrap();
288        loop {
289            n += self.write_level(level)?;
290            if let Some(l) = iter.next() {
291                level = l;
292                n += self.write(b"/")?;
293            } else {
294                break;
295            }
296        }
297        Ok(n)
298    }
299}
300
301impl<W: io::Write + ?Sized> WriteTopicExt for W {}
302
303fn is_system<T: AsRef<str>>(s: T) -> bool {
304    s.as_ref().starts_with('$')
305}
306
307fn recover_bstr(superset: &ByteString, subset: &str) -> ByteString {
308    unsafe {
309        ByteString::from_bytes_unchecked(superset.as_bytes().slice_ref(subset.as_bytes()))
310    }
311}
312
313#[cfg(test)]
314mod tests {
315    use super::*;
316    use test_case::test_case;
317
318    #[test_case("abc" => true; "pass_norm1")]
319    #[test_case("a/b" => true; "pass_norm2")]
320    #[test_case("/" => true; "pass_norm3")]
321    #[test_case("//" => true; "pass_norm4")]
322    #[test_case("a/b/+" => true; "pass_plus1")]
323    #[test_case("+/a" => true; "pass_plus2")]
324    #[test_case("+" => true; "pass_plus3")]
325    #[test_case("+//+" => true; "pass_plus4")]
326    #[test_case("a/b/#" => true; "pass_hash1")]
327    #[test_case("#" => true; "pass_hash2")]
328    #[test_case("/#" => true; "pass_hash3")]
329    #[test_case("++" => false; "fail_plus1")]
330    #[test_case("b+/" => false; "fail_plus2")]
331    #[test_case("a/+b" => false; "fail_plus3")]
332    #[test_case("+#" => false; "fail_hash1")]
333    #[test_case("a#" => false; "fail_hash2")]
334    #[test_case("a/#/" => false; "fail_hash3")]
335    #[test_case("a/#b" => false; "fail_hash4")]
336    #[test_case("a/##" => false; "fail_hash5")]
337    #[test_case("a/#+" => false; "fail_hash6")]
338    fn check_is_valid(topic_filter: &'static str) -> bool {
339        is_valid(topic_filter)
340    }
341
342    fn lvl_normal<T: AsRef<str>>(s: T) -> TopicFilterLevel {
343        assert!(
344            !s.as_ref().contains(['+', '#']),
345            "invalid normal level `{}` contains +|#",
346            s.as_ref()
347        );
348        TopicFilterLevel::Normal(s.as_ref().into())
349    }
350
351    fn lvl_sys<T: AsRef<str>>(s: T) -> TopicFilterLevel {
352        assert!(
353            !s.as_ref().contains(['+', '#']),
354            "invalid normal level `{}` contains +|#",
355            s.as_ref()
356        );
357        assert!(
358            s.as_ref().starts_with('$'),
359            "invalid metadata level `{}` not starts with $",
360            s.as_ref()
361        );
362        TopicFilterLevel::System(s.as_ref().into())
363    }
364
365    fn topic(topic: &'static str) -> TopicFilter {
366        TopicFilter::try_from(ByteString::from_static(topic)).unwrap()
367    }
368
369    #[test_case("level" => Ok(vec![lvl_normal("level")]) ; "1")]
370    #[test_case("level/+" => Ok(vec![lvl_normal("level"), TopicFilterLevel::SingleWildcard]) ; "2")]
371    #[test_case("a//#" => Ok(vec![lvl_normal("a"), TopicFilterLevel::Blank, TopicFilterLevel::MultiWildcard]) ; "3")]
372    #[test_case("$a///#" => Ok(vec![lvl_sys("$a"), TopicFilterLevel::Blank, TopicFilterLevel::Blank, TopicFilterLevel::MultiWildcard]) ; "4")]
373    #[test_case("$a/#/" => Err(TopicFilterError::InvalidTopic) ; "5")]
374    #[test_case("a+b" => Err(TopicFilterError::InvalidLevel) ; "6")]
375    #[test_case("a/+b" => Err(TopicFilterError::InvalidLevel) ; "7")]
376    #[test_case("$a/$b/" => Ok(vec![lvl_sys("$a"), lvl_normal("$b"), TopicFilterLevel::Blank]) ; "8")]
377    #[test_case("#/a" => Err(TopicFilterError::InvalidTopic) ; "10")]
378    #[test_case("" => Err(TopicFilterError::InvalidTopic) ; "11")]
379    #[test_case("/finance" => Ok(vec![TopicFilterLevel::Blank, lvl_normal("finance")]) ; "12")]
380    #[test_case("finance/" => Ok(vec![lvl_normal("finance"), TopicFilterLevel::Blank]) ; "13")]
381    fn parsing(input: &str) -> Result<Vec<TopicFilterLevel>, TopicFilterError> {
382        TopicFilter::try_from(ByteString::from(input)).map(|t| t.levels().to_vec())
383    }
384
385    #[test_case(vec![lvl_normal("sport"), lvl_normal("tennis"), lvl_normal("player1")] => true; "1")]
386    #[test_case(vec![lvl_normal("sport"), lvl_normal("tennis"), TopicFilterLevel::MultiWildcard] => true; "2")]
387    #[test_case(vec![lvl_sys("$SYS"), lvl_normal("tennis"), lvl_normal("player1")] => true; "3")]
388    #[test_case(vec![lvl_normal("sport"), TopicFilterLevel::SingleWildcard, lvl_normal("player1")] => true; "4")]
389    #[test_case(vec![lvl_normal("sport"), TopicFilterLevel::MultiWildcard, lvl_normal("player1")] => false; "5")]
390    #[test_case(vec![lvl_normal("sport"), lvl_sys("$SYS"), lvl_normal("player1")] => false; "6")]
391    fn topic_is_valid(levels: Vec<TopicFilterLevel>) -> bool {
392        TopicFilter::try_from(levels).is_ok()
393    }
394
395    #[test]
396    fn test_multi_wildcard_topic() {
397        assert!(topic("sport/tennis/#").matches_filter(&TopicFilter(vec![
398            lvl_normal("sport"),
399            lvl_normal("tennis"),
400            TopicFilterLevel::MultiWildcard
401        ])));
402
403        assert!(topic("sport/tennis/#").matches_topic("sport/tennis"));
404
405        assert!(topic("#").matches_filter(&TopicFilter(vec![TopicFilterLevel::MultiWildcard])));
406    }
407
408    #[test]
409    fn test_single_wildcard_topic() {
410        assert!(topic("+").matches_filter(
411            &TopicFilter::try_from(vec![TopicFilterLevel::SingleWildcard]).unwrap()
412        ));
413
414        assert!(topic("+/tennis/#").matches_filter(&TopicFilter(vec![
415            TopicFilterLevel::SingleWildcard,
416            lvl_normal("tennis"),
417            TopicFilterLevel::MultiWildcard
418        ])));
419
420        assert!(topic("sport/+/player1").matches_filter(&TopicFilter(vec![
421            lvl_normal("sport"),
422            TopicFilterLevel::SingleWildcard,
423            lvl_normal("player1")
424        ])));
425    }
426
427    #[test]
428    fn test_write_topic() {
429        let mut v = vec![];
430        let t = TopicFilter(vec![
431            TopicFilterLevel::SingleWildcard,
432            lvl_normal("tennis"),
433            TopicFilterLevel::MultiWildcard,
434        ]);
435
436        assert_eq!(v.write_topic(&t).unwrap(), 10);
437        assert_eq!(v, b"+/tennis/#");
438
439        assert_eq!(format!("{t}"), "+/tennis/#");
440        assert_eq!(t.to_string(), "+/tennis/#");
441    }
442
443    #[test_case("test", "test" => true)]
444    #[test_case("$SYS", "$SYS" => true)]
445    #[test_case("sport/tennis/player1/#", "sport/tennis/player1" => true)]
446    #[test_case("sport/tennis/player1/#", "sport/tennis/player1/score" => true)]
447    #[test_case("sport/tennis/player1/#", "sport/tennis/player1/score/wimbledon" => true)]
448    #[test_case("sport/#", "sport" => true)]
449    #[test_case("sport/tennis/+", "sport/tennis/player1" => true)]
450    #[test_case("sport/tennis/+", "sport/tennis/player2" => true)]
451    #[test_case("sport/tennis/+", "sport/tennis/player1/ranking" => false)]
452    #[test_case("sport/+", "sport" => false; "single1")]
453    #[test_case("sport/+", "sport/" => true; "single2")]
454    #[test_case("+/+", "/finance" => true; "single3")]
455    #[test_case("/+", "/finance" => true; "single4")]
456    #[test_case("+", "/finance" => false; "single5")]
457    #[test_case("#", "$SYS" => false; "sys1")]
458    #[test_case("+/monitor/Clients", "$SYS/monitor/Clients" => false; "sys2")]
459    #[test_case("$SYS/#", "$SYS/" => true; "sys3")]
460    #[test_case("$SYS/monitor/+", "$SYS/monitor/Clients" => true; "sys4")]
461    #[test_case("#", "/$SYS/monitor/Clients" => true; "sys5")]
462    #[test_case("+", "$SYS" => false; "sys6")]
463    #[test_case("+/#", "$SYS" => false; "sys7")]
464    fn matches_topic(filter: &'static str, topic_str: &'static str) -> bool {
465        topic(filter).matches_topic(topic_str)
466    }
467
468    #[test_case("a/b", "a/b" => true; "1")]
469    #[test_case("a/b", "a/+" => false; "2")]
470    #[test_case("a/b", "a/#" => false; "3")]
471    #[test_case("a/+", "a/#" => false; "4")]
472    #[test_case("a/+", "a/b" => true; "5")]
473    #[test_case("+/+", "/" => true; "6")]
474    #[test_case("+/+", "#" => false; "7")]
475    #[test_case("+", "#" => false; "8")]
476    #[test_case("#", "+" => true; "9")]
477    #[test_case("#", "#" => true; "10")]
478    #[test_case("a/#", "a/+/+" => true; "11")]
479    #[test_case("a/+/normal/+", "a/$not_sys/normal/+" => true; "12")]
480    #[test_case("a/+/#", "a/b" => true; "13")]
481    fn matches_filter(superset_filter: &'static str, subset_filter: &'static str) -> bool {
482        topic(superset_filter).matches_filter(&topic(subset_filter))
483    }
484}