Skip to main content

cast/cmd/
tempo_policy_args.rs

1use alloy_primitives::{Address, hex};
2use foundry_common::abi::get_func;
3use tempo_contracts::precompiles::IAccountKeychain::{CallScope, SelectorRule};
4
5// Shared Tempo policy flag grammar used by both `cast keychain` and
6// `cast wallet session`. Keeping it here avoids duplicating parsing behavior
7// or making wallet-session commands depend on the larger keychain command module.
8
9/// Parsed selector argument used by policy-editing commands.
10#[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
19/// Parse a selector string into 4-byte selector bytes.
20///
21/// Accepts 4-byte hex (`0xd09de08a`), a full signature
22/// (`transfer(address,uint256)`), or a well-known TIP-20 shorthand.
23pub(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
55/// Parse a selector string into a named selector argument.
56pub(crate) fn parse_selector_arg(s: &str) -> Result<SelectorArg, String> {
57    parse_selector_bytes(s).map(SelectorArg)
58}
59
60/// Parse a `TARGET[:SELECTORS[@RECIPIENTS]]` scope string.
61pub(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
135/// Parse a period string like `10m`, `7d`, or `3600s`.
136pub(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}