Skip to main content

r402_protocol/network/
pattern.rs

1//! Chain-ID patterns: exact, wildcard, and set.
2
3use std::collections::HashSet;
4use std::fmt;
5use std::str::FromStr;
6
7use serde::{Deserialize, Deserializer, Serialize, Serializer, de};
8
9use super::id::{ChainId, ChainIdFormatError};
10
11/// Pattern for matching [`ChainId`] values.
12///
13/// - Wildcard: `"eip155:*"`
14/// - Exact: `"eip155:8453"`
15/// - Set: `"eip155:{1,8453,137}"`
16#[derive(Debug, Clone, PartialEq, Eq)]
17#[non_exhaustive]
18pub enum ChainIdPattern {
19    /// Any chain in `namespace`.
20    Wildcard {
21        /// Namespace to match (`eip155`, `solana`, …).
22        namespace: String,
23    },
24    /// One specific chain.
25    Exact {
26        /// Namespace of the chain.
27        namespace: String,
28        /// Reference of the chain.
29        reference: String,
30    },
31    /// Any reference in `references` within `namespace`.
32    Set {
33        /// Namespace of the chains.
34        namespace: String,
35        /// Set of chain references to match.
36        references: HashSet<String>,
37    },
38}
39
40impl ChainIdPattern {
41    /// Wildcard matching any chain in `namespace`.
42    pub fn wildcard<S: Into<String>>(namespace: S) -> Self {
43        Self::Wildcard {
44            namespace: namespace.into(),
45        }
46    }
47
48    /// Exact match on namespace and reference.
49    pub fn exact<N: Into<String>, R: Into<String>>(namespace: N, reference: R) -> Self {
50        Self::Exact {
51            namespace: namespace.into(),
52            reference: reference.into(),
53        }
54    }
55
56    /// Set match on any of `references` in `namespace`.
57    pub fn set<N: Into<String>>(namespace: N, references: HashSet<String>) -> Self {
58        Self::Set {
59            namespace: namespace.into(),
60            references,
61        }
62    }
63
64    /// Whether `chain_id` matches this pattern.
65    #[must_use]
66    pub fn matches(&self, chain_id: &ChainId) -> bool {
67        match self {
68            Self::Wildcard { namespace } => chain_id.namespace() == namespace,
69            Self::Exact {
70                namespace,
71                reference,
72            } => chain_id.namespace() == namespace && chain_id.reference() == reference,
73            Self::Set {
74                namespace,
75                references,
76            } => chain_id.namespace() == namespace && references.contains(chain_id.reference()),
77        }
78    }
79
80    /// Namespace of this pattern.
81    #[must_use]
82    pub fn namespace(&self) -> &str {
83        match self {
84            Self::Wildcard { namespace }
85            | Self::Exact { namespace, .. }
86            | Self::Set { namespace, .. } => namespace,
87        }
88    }
89}
90
91impl fmt::Display for ChainIdPattern {
92    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
93        match self {
94            Self::Wildcard { namespace } => write!(f, "{namespace}:*"),
95            Self::Exact {
96                namespace,
97                reference,
98            } => write!(f, "{namespace}:{reference}"),
99            Self::Set {
100                namespace,
101                references,
102            } => {
103                let refs: Vec<&str> = references.iter().map(AsRef::as_ref).collect();
104                write!(f, "{}:{{{}}}", namespace, refs.join(","))
105            }
106        }
107    }
108}
109
110impl FromStr for ChainIdPattern {
111    type Err = ChainIdFormatError;
112
113    fn from_str(s: &str) -> Result<Self, Self::Err> {
114        let (namespace, rest) = s
115            .split_once(':')
116            .ok_or_else(|| ChainIdFormatError(s.into()))?;
117
118        if namespace.is_empty() {
119            return Err(ChainIdFormatError(s.into()));
120        }
121
122        if rest == "*" {
123            return Ok(Self::wildcard(namespace));
124        }
125
126        if let Some(inner) = rest.strip_prefix('{').and_then(|r| r.strip_suffix('}')) {
127            let items: Vec<&str> = inner.split(',').map(str::trim).collect();
128            if items.is_empty() || items.iter().any(|item| item.is_empty()) {
129                return Err(ChainIdFormatError(s.into()));
130            }
131            let references = items.into_iter().map(Into::into).collect();
132            return Ok(Self::set(namespace, references));
133        }
134
135        if rest.is_empty() {
136            return Err(ChainIdFormatError(s.into()));
137        }
138
139        Ok(Self::exact(namespace, rest))
140    }
141}
142
143impl Serialize for ChainIdPattern {
144    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
145    where
146        S: Serializer,
147    {
148        serializer.serialize_str(&self.to_string())
149    }
150}
151
152impl<'de> Deserialize<'de> for ChainIdPattern {
153    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
154    where
155        D: Deserializer<'de>,
156    {
157        let s = String::deserialize(deserializer)?;
158        Self::from_str(&s).map_err(de::Error::custom)
159    }
160}
161
162impl From<ChainId> for ChainIdPattern {
163    fn from(chain_id: ChainId) -> Self {
164        let (namespace, reference) = chain_id.into_parts();
165        Self::exact(namespace, reference)
166    }
167}
168
169#[cfg(test)]
170mod tests {
171    use super::*;
172
173    #[test]
174    fn wildcard_matches() {
175        let pattern = ChainIdPattern::wildcard("eip155");
176        assert!(pattern.matches(&ChainId::new("eip155", "1")));
177        assert!(pattern.matches(&ChainId::new("eip155", "8453")));
178        assert!(pattern.matches(&ChainId::new("eip155", "137")));
179        assert!(!pattern.matches(&ChainId::new("solana", "mainnet")));
180    }
181
182    #[test]
183    fn exact_matches() {
184        let pattern = ChainIdPattern::exact("eip155", "1");
185        assert!(pattern.matches(&ChainId::new("eip155", "1")));
186        assert!(!pattern.matches(&ChainId::new("eip155", "8453")));
187        assert!(!pattern.matches(&ChainId::new("solana", "1")));
188    }
189
190    #[test]
191    fn set_matches() {
192        let references: HashSet<String> =
193            ["1", "8453", "137"].into_iter().map(String::from).collect();
194        let pattern = ChainIdPattern::set("eip155", references);
195        assert!(pattern.matches(&ChainId::new("eip155", "1")));
196        assert!(pattern.matches(&ChainId::new("eip155", "8453")));
197        assert!(pattern.matches(&ChainId::new("eip155", "137")));
198        assert!(!pattern.matches(&ChainId::new("eip155", "42")));
199        assert!(!pattern.matches(&ChainId::new("solana", "1")));
200    }
201
202    #[test]
203    fn namespace_accessor() {
204        let wildcard = ChainIdPattern::wildcard("eip155");
205        assert_eq!(wildcard.namespace(), "eip155");
206
207        let exact = ChainIdPattern::exact("solana", "mainnet");
208        assert_eq!(exact.namespace(), "solana");
209
210        let references: HashSet<String> = std::iter::once("1").map(String::from).collect();
211        let set = ChainIdPattern::set("eip155", references);
212        assert_eq!(set.namespace(), "eip155");
213    }
214}