1use alloy_primitives::{Address, hex};
2use foundry_common::abi::get_func;
3use tempo_contracts::precompiles::IAccountKeychain::{CallScope, SelectorRule};
4
5#[derive(Debug, Clone, Copy)]
11pub struct SelectorArg([u8; 4]);
12
13impl SelectorArg {
14 pub(crate) const fn into_bytes(self) -> [u8; 4] {
15 self.0
16 }
17}
18
19pub(crate) fn parse_selector_bytes(s: &str) -> Result<[u8; 4], String> {
24 let s = s.trim();
25 if s.starts_with("0x") || s.starts_with("0X") {
26 let hex_str = &s[2..];
27 if hex_str.len() != 8 {
28 return Err(format!("hex selector must be 4 bytes (8 hex chars), got: {s}"));
29 }
30 let bytes = hex::decode(hex_str).map_err(|e| format!("invalid hex selector '{s}': {e}"))?;
31 let mut arr = [0u8; 4];
32 arr.copy_from_slice(&bytes);
33 Ok(arr)
34 } else {
35 let sig = if s.contains('(') || s.contains(')') {
36 s.to_string()
37 } else {
38 match s {
39 "transfer" => "transfer(address,uint256)".to_string(),
40 "approve" => "approve(address,uint256)".to_string(),
41 "transferFrom" => "transferFrom(address,address,uint256)".to_string(),
42 "transferWithMemo" => "transferWithMemo(address,uint256,bytes32)".to_string(),
43 "transferFromWithMemo" => {
44 "transferFromWithMemo(address,address,uint256,bytes32)".to_string()
45 }
46 _ => format!("{s}()"),
47 }
48 };
49 get_func(&sig)
50 .map(|func| func.selector().into())
51 .map_err(|e| format!("invalid function signature '{sig}': {e}"))
52 }
53}
54
55pub(crate) fn parse_selector_arg(s: &str) -> Result<SelectorArg, String> {
57 parse_selector_bytes(s).map(SelectorArg)
58}
59
60pub(crate) fn parse_scope(s: &str) -> Result<CallScope, String> {
62 let (target_str, selectors_str) = match s.split_once(':') {
63 Some((t, sel)) => (t, Some(sel)),
64 None => (s, None),
65 };
66
67 let target: Address =
68 target_str.parse().map_err(|e| format!("invalid target address '{target_str}': {e}"))?;
69
70 let selector_rules = match selectors_str {
71 None => vec![],
72 Some(sel_str) => parse_selector_rules(sel_str)?,
73 };
74
75 Ok(CallScope { target, selectorRules: selector_rules })
76}
77
78fn parse_selector_rules(s: &str) -> Result<Vec<SelectorRule>, String> {
79 let mut rules = Vec::new();
80
81 for part in split_selector_rule_parts(s) {
82 let part = part.trim();
83 if part.is_empty() {
84 continue;
85 }
86
87 let (selector_str, recipients_str) = match part.split_once('@') {
88 Some((sel, recip)) => (sel, Some(recip)),
89 None => (part, None),
90 };
91
92 let selector = parse_selector_bytes(selector_str)?;
93
94 let recipients = match recipients_str {
95 None => vec![],
96 Some(r) => r
97 .split(',')
98 .filter(|s| !s.trim().is_empty())
99 .map(|addr_str| {
100 let addr_str = addr_str.trim();
101 addr_str
102 .parse::<Address>()
103 .map_err(|e| format!("invalid recipient address '{addr_str}': {e}"))
104 })
105 .collect::<Result<Vec<_>, _>>()?,
106 };
107
108 rules.push(SelectorRule { selector: selector.into(), recipients });
109 }
110
111 Ok(rules)
112}
113
114fn split_selector_rule_parts(s: &str) -> Vec<&str> {
115 let mut parts = Vec::new();
116 let mut depth = 0usize;
117 let mut start = 0usize;
118
119 for (idx, ch) in s.char_indices() {
120 match ch {
121 '(' => depth += 1,
122 ')' => depth = depth.saturating_sub(1),
123 ',' if depth == 0 => {
124 parts.push(&s[start..idx]);
125 start = idx + ch.len_utf8();
126 }
127 _ => {}
128 }
129 }
130
131 parts.push(&s[start..]);
132 parts
133}
134
135pub(crate) fn parse_period(s: &str) -> Result<u64, String> {
137 let s = s.trim();
138 if s.is_empty() {
139 return Err("period cannot be empty".to_string());
140 }
141
142 let split = s.find(|c: char| !c.is_ascii_digit()).unwrap_or(s.len());
143 if split == 0 {
144 return Err(format!(
145 "invalid period '{s}': expected a number followed by s, m, h, d, or w"
146 ));
147 }
148
149 let value: u64 =
150 s[..split].parse().map_err(|e| format!("invalid period value '{}': {e}", &s[..split]))?;
151 let multiplier = match &s[split..].to_ascii_lowercase()[..] {
152 "" | "s" => 1,
153 "m" => 60,
154 "h" => 60 * 60,
155 "d" => 24 * 60 * 60,
156 "w" => 7 * 24 * 60 * 60,
157 unit => {
158 return Err(format!(
159 "invalid period unit '{unit}' in '{s}' (expected s, m, h, d, or w)"
160 ));
161 }
162 };
163
164 value.checked_mul(multiplier).ok_or_else(|| format!("period '{s}' is too large"))
165}
166
167#[cfg(test)]
168mod tests {
169 use super::*;
170 use alloy_primitives::keccak256;
171 use std::str::FromStr;
172
173 #[test]
174 fn parse_selector_bytes_named() {
175 let sel = parse_selector_bytes("transfer").unwrap();
176 assert_eq!(sel, keccak256(b"transfer(address,uint256)")[..4]);
177
178 let sel = parse_selector_bytes("approve").unwrap();
179 assert_eq!(sel, keccak256(b"approve(address,uint256)")[..4]);
180
181 let sel = parse_selector_bytes("transferWithMemo").unwrap();
182 assert_eq!(sel, keccak256(b"transferWithMemo(address,uint256,bytes32)")[..4]);
183 }
184
185 #[test]
186 fn parse_selector_bytes_hex() {
187 let sel = parse_selector_bytes("0xaabbccdd").unwrap();
188 assert_eq!(sel, [0xaa, 0xbb, 0xcc, 0xdd]);
189
190 let sel = parse_selector_bytes("0xd09de08a").unwrap();
191 assert_eq!(sel, [0xd0, 0x9d, 0xe0, 0x8a]);
192 }
193
194 #[test]
195 fn parse_selector_bytes_hex_invalid() {
196 assert!(parse_selector_bytes("0xaabb").is_err());
197 assert!(parse_selector_bytes("0xaabbccddee").is_err());
198 assert!(parse_selector_bytes("0xzzzzzzzz").is_err());
199 }
200
201 #[test]
202 fn parse_selector_bytes_full_signature() {
203 let sel = parse_selector_bytes("increment()").unwrap();
204 assert_eq!(sel, keccak256(b"increment()")[..4]);
205
206 let sel = parse_selector_bytes("transfer(address,uint256)").unwrap();
207 assert_eq!(sel, keccak256(b"transfer(address,uint256)")[..4]);
208 }
209
210 #[test]
211 fn parse_selector_bytes_rejects_invalid_signature() {
212 assert!(parse_selector_bytes("").is_err());
213 assert!(parse_selector_bytes("transfer(address,uint256").is_err());
214 assert!(parse_selector_bytes("transfer)").is_err());
215 }
216
217 #[test]
218 fn parse_scope_hex_selector_with_recipient() {
219 let scope = parse_scope(
220 "0x20c0000000000000000000000000000000000001:0xaabbccdd@0x1111111111111111111111111111111111111111",
221 )
222 .unwrap();
223 assert_eq!(scope.selectorRules.len(), 1);
224 assert_eq!(scope.selectorRules[0].selector.0, [0xaa, 0xbb, 0xcc, 0xdd]);
225 assert_eq!(scope.selectorRules[0].recipients.len(), 1);
226 }
227
228 #[test]
229 fn parse_scope_target_only() {
230 let scope = parse_scope("0x86A2EE8FAf9A840F7a2c64CA3d51209F9A02081D").unwrap();
231 assert_eq!(
232 scope.target,
233 Address::from_str("0x86A2EE8FAf9A840F7a2c64CA3d51209F9A02081D").unwrap()
234 );
235 assert!(scope.selectorRules.is_empty());
236 }
237
238 #[test]
239 fn parse_scope_with_selectors() {
240 let scope =
241 parse_scope("0x20c0000000000000000000000000000000000001:transfer,approve").unwrap();
242 assert_eq!(scope.selectorRules.len(), 2);
243 assert!(scope.selectorRules[0].recipients.is_empty());
244 assert!(scope.selectorRules[1].recipients.is_empty());
245 }
246
247 #[test]
248 fn parse_scope_hex_selector() {
249 let scope = parse_scope("0x86A2EE8FAf9A840F7a2c64CA3d51209F9A02081D:0xaabbccdd").unwrap();
250 assert_eq!(scope.selectorRules.len(), 1);
251 assert_eq!(scope.selectorRules[0].selector.0, [0xaa, 0xbb, 0xcc, 0xdd]);
252 assert!(scope.selectorRules[0].recipients.is_empty());
253 }
254
255 #[test]
256 fn parse_scope_selector_with_recipient() {
257 let scope = parse_scope(
258 "0x20c0000000000000000000000000000000000001:transfer@0x1111111111111111111111111111111111111111",
259 )
260 .unwrap();
261 assert_eq!(scope.selectorRules.len(), 1);
262 assert_eq!(scope.selectorRules[0].recipients.len(), 1);
263 }
264
265 #[test]
266 fn parse_scope_full_signatures_split_outside_parentheses() {
267 let scope = parse_scope(
268 "0x20c0000000000000000000000000000000000001:transfer(address,uint256),approve(address,uint256)",
269 )
270 .unwrap();
271 assert_eq!(scope.selectorRules.len(), 2);
272 assert_eq!(scope.selectorRules[0].selector.0, keccak256(b"transfer(address,uint256)")[..4]);
273 assert_eq!(scope.selectorRules[1].selector.0, keccak256(b"approve(address,uint256)")[..4]);
274 }
275
276 #[test]
277 fn parse_period_units() {
278 assert_eq!(parse_period("0").unwrap(), 0);
279 assert_eq!(parse_period("30s").unwrap(), 30);
280 assert_eq!(parse_period("5m").unwrap(), 300);
281 assert_eq!(parse_period("2h").unwrap(), 7200);
282 assert_eq!(parse_period("7d").unwrap(), 604800);
283 assert_eq!(parse_period("2w").unwrap(), 1209600);
284 assert!(parse_period("1mo").is_err());
285 }
286}