Skip to main content

relay_pattern/
typed.rs

1use core::fmt;
2use std::marker::PhantomData;
3use std::ops::Deref;
4
5use crate::{Error, Pattern, Patterns, PatternsBuilderConfigured};
6
7/// Compile time configuration for a [`TypedPattern`].
8pub trait PatternConfig {
9    /// Configures the pattern to match case insensitive.
10    const CASE_INSENSITIVE: bool = false;
11    /// Configures the maximum allowed complexity of the pattern.
12    const MAX_COMPLEXITY: u64 = u64::MAX;
13}
14
15/// The default pattern.
16///
17/// Equivalent to [`Pattern::new`].
18pub struct DefaultPatternConfig;
19
20impl PatternConfig for DefaultPatternConfig {}
21
22/// The default pattern but with case insensitive matching.
23///
24/// See: [`crate::PatternBuilder::case_insensitive`].
25pub struct CaseInsensitive;
26
27impl PatternConfig for CaseInsensitive {
28    const CASE_INSENSITIVE: bool = true;
29}
30
31/// A [`Pattern`] with compile time encoded [`PatternConfig`].
32///
33/// Encoding the pattern configuration allows context dependent serialization
34/// and usage of patterns and ensures a consistent usage of configuration options
35/// throught the code.
36///
37/// Often repeated configuration can be grouped into custom and importable configurations.
38///
39/// ```
40/// struct MetricConfig;
41///
42/// impl relay_pattern::PatternConfig for MetricConfig {
43///     const CASE_INSENSITIVE: bool = false;
44///     // More configuration ...
45/// }
46///
47/// type MetricPattern = relay_pattern::TypedPattern<MetricConfig>;
48///
49/// let pattern = MetricPattern::new("[cd]:foo/bar").unwrap();
50/// assert!(pattern.is_match("c:foo/bar"));
51/// ```
52pub struct TypedPattern<C = DefaultPatternConfig> {
53    pattern: Pattern,
54    _phantom: PhantomData<C>,
55}
56
57impl<C: PatternConfig> TypedPattern<C> {
58    /// Creates a new [`TypedPattern`] using the provided pattern and config `C`.
59    ///
60    /// ```
61    /// use relay_pattern::{Pattern, TypedPattern, CaseInsensitive};
62    ///
63    /// let pattern = TypedPattern::<CaseInsensitive>::new("foo*").unwrap();
64    /// assert!(pattern.is_match("FOOBAR"));
65    ///
66    /// // Equivalent to:
67    /// let pattern = Pattern::builder("foo*").case_insensitive(true).build().unwrap();
68    /// assert!(pattern.is_match("FOOBAR"));
69    /// ```
70    pub fn new(pattern: &str) -> Result<Self, Error> {
71        Pattern::builder(pattern)
72            .case_insensitive(C::CASE_INSENSITIVE)
73            .max_complexity(C::MAX_COMPLEXITY)
74            .build()
75            .map(|pattern| Self {
76                pattern,
77                _phantom: PhantomData,
78            })
79    }
80}
81
82#[cfg(feature = "serde")]
83impl<'de, C: PatternConfig> serde::Deserialize<'de> for TypedPattern<C> {
84    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
85    where
86        D: serde::Deserializer<'de>,
87    {
88        let pattern = <std::borrow::Cow<'_, str>>::deserialize(deserializer)?;
89        Self::new(&pattern).map_err(serde::de::Error::custom)
90    }
91}
92
93#[cfg(feature = "serde")]
94impl<C> serde::Serialize for TypedPattern<C> {
95    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
96    where
97        S: serde::Serializer,
98    {
99        serializer.collect_str(&self.pattern)
100    }
101}
102
103impl<C> PartialEq for TypedPattern<C> {
104    fn eq(&self, other: &Self) -> bool {
105        self.pattern.eq(other)
106    }
107}
108
109impl<C> Eq for TypedPattern<C> {}
110
111impl<C> From<TypedPattern<C>> for Pattern {
112    fn from(value: TypedPattern<C>) -> Self {
113        value.pattern
114    }
115}
116
117impl<C> AsRef<Pattern> for TypedPattern<C> {
118    fn as_ref(&self) -> &Pattern {
119        &self.pattern
120    }
121}
122
123impl<C> Deref for TypedPattern<C> {
124    type Target = Pattern;
125
126    fn deref(&self) -> &Self::Target {
127        &self.pattern
128    }
129}
130
131impl<C> fmt::Debug for TypedPattern<C> {
132    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
133        self.pattern.fmt(f)
134    }
135}
136
137/// [`Patterns`] with a compile time configured [`PatternConfig`].
138pub struct TypedPatterns<C = DefaultPatternConfig> {
139    patterns: Patterns,
140    _phantom: PhantomData<C>,
141}
142
143impl<C: PatternConfig> TypedPatterns<C> {
144    pub fn builder() -> TypedPatternsBuilder<C> {
145        let builder = Patterns::builder()
146            .case_insensitive(C::CASE_INSENSITIVE)
147            .patterns();
148
149        TypedPatternsBuilder {
150            builder,
151            _phantom: PhantomData,
152        }
153    }
154}
155
156impl<C: PatternConfig> Default for TypedPatterns<C> {
157    fn default() -> Self {
158        Self::builder().build()
159    }
160}
161
162impl<C> PartialEq for TypedPatterns<C> {
163    fn eq(&self, other: &Self) -> bool {
164        self.patterns.eq(other)
165    }
166}
167
168impl<C> Eq for TypedPatterns<C> {}
169
170impl<C: PatternConfig> From<String> for TypedPatterns<C> {
171    fn from(value: String) -> Self {
172        [value].into_iter().collect()
173    }
174}
175
176impl<C: PatternConfig> From<Vec<String>> for TypedPatterns<C> {
177    fn from(value: Vec<String>) -> Self {
178        value.into_iter().collect()
179    }
180}
181
182impl<C: PatternConfig, const N: usize> From<[String; N]> for TypedPatterns<C> {
183    fn from(value: [String; N]) -> Self {
184        value.into_iter().collect()
185    }
186}
187
188/// Creates [`Patterns`] from an iterator of strings.
189///
190/// Invalid patterns are ignored.
191impl<C: PatternConfig> FromIterator<String> for TypedPatterns<C> {
192    fn from_iter<T: IntoIterator<Item = String>>(iter: T) -> Self {
193        let mut builder = Self::builder();
194        for pattern in iter.into_iter() {
195            let _err = builder.add(pattern);
196            #[cfg(debug_assertions)]
197            _err.expect("all patterns should be valid patterns");
198        }
199        builder.build()
200    }
201}
202
203impl<C> From<TypedPatterns<C>> for Patterns {
204    fn from(value: TypedPatterns<C>) -> Self {
205        value.patterns
206    }
207}
208
209impl<C> AsRef<Patterns> for TypedPatterns<C> {
210    fn as_ref(&self) -> &Patterns {
211        &self.patterns
212    }
213}
214
215impl<C> Deref for TypedPatterns<C> {
216    type Target = Patterns;
217
218    fn deref(&self) -> &Self::Target {
219        &self.patterns
220    }
221}
222
223impl<C> Clone for TypedPatterns<C> {
224    fn clone(&self) -> Self {
225        Self {
226            patterns: self.patterns.clone(),
227            _phantom: PhantomData,
228        }
229    }
230}
231
232impl<C> fmt::Debug for TypedPatterns<C> {
233    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
234        self.patterns.fmt(f)
235    }
236}
237
238/// Deserializes patterns from a sequence of strings.
239///
240/// Invalid patterns are ignored while deserializing.
241#[cfg(feature = "serde")]
242impl<'de, C: PatternConfig> serde::Deserialize<'de> for TypedPatterns<C> {
243    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
244    where
245        D: serde::Deserializer<'de>,
246    {
247        struct Visitor<C>(PhantomData<C>);
248
249        impl<'a, C: PatternConfig> serde::de::Visitor<'a> for Visitor<C> {
250            type Value = TypedPatterns<C>;
251
252            fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
253                formatter.write_str("a sequence of patterns")
254            }
255
256            fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
257            where
258                A: serde::de::SeqAccess<'a>,
259            {
260                let mut builder = TypedPatterns::<C>::builder();
261
262                while let Some(item) = seq.next_element()? {
263                    // Ignore invalid patterns as documented.
264                    let _err = builder.add(item);
265                    #[cfg(debug_assertions)]
266                    _err.expect("all patterns should be valid patterns");
267                }
268
269                Ok(builder.build())
270            }
271        }
272
273        deserializer.deserialize_seq(Visitor(PhantomData))
274    }
275}
276
277#[cfg(feature = "serde")]
278impl<C> serde::Serialize for TypedPatterns<C> {
279    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
280    where
281        S: serde::Serializer,
282    {
283        self.patterns.serialize(serializer)
284    }
285}
286
287pub struct TypedPatternsBuilder<C> {
288    builder: PatternsBuilderConfigured,
289    _phantom: PhantomData<C>,
290}
291
292impl<C: PatternConfig> TypedPatternsBuilder<C> {
293    /// Adds a pattern to the builder.
294    pub fn add(&mut self, pattern: String) -> Result<&mut Self, Error> {
295        self.builder.add(&pattern)?;
296        Ok(self)
297    }
298
299    /// Builds a [`TypedPatterns`] from the contained patterns.
300    pub fn build(self) -> TypedPatterns<C> {
301        TypedPatterns {
302            patterns: self.builder.build(),
303            _phantom: PhantomData,
304        }
305    }
306
307    /// Builds a [`TypedPatterns`] from the contained patterns and clears the builder.
308    pub fn take(&mut self) -> TypedPatterns<C> {
309        TypedPatterns {
310            patterns: self.builder.take(),
311            _phantom: PhantomData,
312        }
313    }
314}
315
316#[cfg(test)]
317mod tests {
318    use super::*;
319
320    #[test]
321    fn test_pattern_default() {
322        let pattern: TypedPattern = TypedPattern::new("*[rt]x").unwrap();
323        assert!(pattern.is_match("f/o_rx"));
324        assert!(pattern.is_match("f/o_tx"));
325        assert!(pattern.is_match("F/o_tx"));
326        // case sensitive
327        assert!(!pattern.is_match("f/o_Tx"));
328        assert!(!pattern.is_match("f/o_rX"));
329    }
330
331    #[test]
332    fn test_pattern_case_insensitive() {
333        let pattern: TypedPattern<CaseInsensitive> = TypedPattern::new("*[rt]x").unwrap();
334        // case insensitive
335        assert!(pattern.is_match("f/o_Tx"));
336        assert!(pattern.is_match("f/o_rX"));
337    }
338
339    #[test]
340    fn test_pattern_eq() {
341        let pattern1: TypedPattern<CaseInsensitive> = TypedPattern::new("Foo**").unwrap();
342        let pattern2: TypedPattern<CaseInsensitive> = TypedPattern::new("Foo**").unwrap();
343        let pattern3: TypedPattern<CaseInsensitive> = TypedPattern::new("foo*").unwrap();
344        assert_eq!(&pattern1, &pattern1);
345        assert_eq!(&pattern1, &pattern2);
346        assert_eq!(&pattern1, &pattern3);
347        assert_eq!(&pattern2, &pattern3);
348    }
349
350    #[test]
351    fn test_pattern_neq() {
352        let pattern1: TypedPattern = TypedPattern::new("Foo**").unwrap();
353        let pattern2: TypedPattern = TypedPattern::new("foo*").unwrap();
354        assert_ne!(&pattern1, &pattern2);
355    }
356
357    #[test]
358    #[cfg(feature = "serde")]
359    fn test_pattern_deserialize() {
360        let pattern: TypedPattern<CaseInsensitive> = serde_json::from_str(r#""*[rt]x""#).unwrap();
361        assert!(pattern.is_match("foobar_rx"));
362    }
363
364    #[test]
365    #[cfg(feature = "serde")]
366    fn test_deserialize_err() {
367        let r: Result<TypedPattern<CaseInsensitive>, _> = serde_json::from_str(r#""[invalid""#);
368        assert!(r.is_err());
369    }
370
371    #[test]
372    #[cfg(feature = "serde")]
373    fn test_pattern_deserialize_complexity() {
374        struct Test;
375        impl PatternConfig for Test {
376            const MAX_COMPLEXITY: u64 = 2;
377        }
378        let r: Result<TypedPattern<Test>, _> = serde_json::from_str(r#""{foo,bar}""#);
379        assert!(r.is_ok());
380        let r: Result<TypedPattern<Test>, _> = serde_json::from_str(r#""{foo,bar,baz}""#);
381        assert!(r.is_err());
382    }
383
384    #[test]
385    #[cfg(feature = "serde")]
386    fn test_pattern_serialize() {
387        let pattern: TypedPattern = TypedPattern::new("*[rt]x").unwrap();
388        assert_eq!(serde_json::to_string(&pattern).unwrap(), r#""*[rt]x""#);
389        let pattern: TypedPattern<CaseInsensitive> = TypedPattern::new("*[rt]x").unwrap();
390        assert_eq!(serde_json::to_string(&pattern).unwrap(), r#""*[rt]x""#);
391    }
392
393    #[test]
394    fn test_patterns_default() {
395        let patterns: TypedPatterns = TypedPatterns::builder()
396            .add("*[rt]x".to_owned())
397            .unwrap()
398            .add("foobar".to_owned())
399            .unwrap()
400            .take();
401        assert!(patterns.is_match("f/o_rx"));
402        assert!(patterns.is_match("foobar"));
403        assert!(!patterns.is_match("Foobar"));
404    }
405
406    #[test]
407    fn test_patterns_case_insensitive() {
408        let patterns: TypedPatterns<CaseInsensitive> = TypedPatterns::builder()
409            .add("*[rt]x".to_owned())
410            .unwrap()
411            .add("foobar".to_owned())
412            .unwrap()
413            .take();
414        assert!(patterns.is_match("f/o_rx"));
415        assert!(patterns.is_match("f/o_Rx"));
416        assert!(patterns.is_match("foobar"));
417        assert!(patterns.is_match("Foobar"));
418    }
419
420    #[test]
421    #[cfg(feature = "serde")]
422    fn test_patterns_deserialize() {
423        let pattern: TypedPatterns<CaseInsensitive> =
424            serde_json::from_str(r#"["*[rt]x","foobar"]"#).unwrap();
425        assert!(pattern.is_match("foobar_rx"));
426        assert!(pattern.is_match("FOOBAR"));
427    }
428
429    #[test]
430    #[cfg(all(feature = "serde", not(debug_assertions)))]
431    fn test_patterns_deserialize_err() {
432        let r: TypedPatterns<CaseInsensitive> =
433            serde_json::from_str(r#"["[invalid","foobar"]"#).unwrap();
434        assert!(r.is_match("foobar"));
435        assert!(r.is_match("FOOBAR"));
436
437        // The invalid element is dropped.
438        assert_eq!(serde_json::to_string(&r).unwrap(), r#"["foobar"]"#);
439    }
440
441    #[test]
442    #[cfg(feature = "serde")]
443    fn test_patterns_serialize() {
444        let pattern: TypedPatterns = TypedPatterns::builder()
445            .add("*[rt]x".to_owned())
446            .unwrap()
447            .add("foobar".to_owned())
448            .unwrap()
449            .take();
450        assert_eq!(
451            serde_json::to_string(&pattern).unwrap(),
452            r#"["*[rt]x","foobar"]"#
453        );
454
455        let pattern: TypedPatterns<CaseInsensitive> = TypedPatterns::builder()
456            .add("*[rt]x".to_owned())
457            .unwrap()
458            .add("foobar".to_owned())
459            .unwrap()
460            .take();
461        assert_eq!(
462            serde_json::to_string(&pattern).unwrap(),
463            r#"["*[rt]x","foobar"]"#
464        );
465    }
466}