r402_protocol/network/
pattern.rs1use 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#[derive(Debug, Clone, PartialEq, Eq)]
17#[non_exhaustive]
18pub enum ChainIdPattern {
19 Wildcard {
21 namespace: String,
23 },
24 Exact {
26 namespace: String,
28 reference: String,
30 },
31 Set {
33 namespace: String,
35 references: HashSet<String>,
37 },
38}
39
40impl ChainIdPattern {
41 pub fn wildcard<S: Into<String>>(namespace: S) -> Self {
43 Self::Wildcard {
44 namespace: namespace.into(),
45 }
46 }
47
48 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 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 #[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 #[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}