Skip to main content

forest/key_management/
wallet.rs

1// Copyright 2019-2026 ChainSafe Systems
2// SPDX-License-Identifier: Apache-2.0, MIT
3
4use std::{convert::TryFrom, str::FromStr};
5
6use super::{KeyInfo, KeyStore, errors::Error, wallet_helpers};
7use crate::shim::{address::Address, crypto::SignatureType};
8use serde::{Deserialize, Serialize};
9
10#[cfg(test)]
11use {
12    crate::shim::crypto::Signature,
13    ahash::{HashMap, HashMapExt as _},
14};
15
16/// A key, this contains a `KeyInfo`, an address, and a public key.
17#[derive(Clone, PartialEq, Debug, Eq, Serialize, Deserialize)]
18pub struct Key {
19    pub key_info: KeyInfo,
20    // Vec<u8> is used because The public keys for BLS and SECP256K1 are not of the same type
21    pub public_key: Vec<u8>,
22    pub address: Address,
23}
24
25impl TryFrom<KeyInfo> for Key {
26    type Error = crate::key_management::errors::Error;
27
28    fn try_from(key_info: KeyInfo) -> Result<Self, Self::Error> {
29        let public_key = wallet_helpers::to_uncompressed_public_key(
30            *key_info.key_type(),
31            key_info.private_key(),
32        )?;
33        let address = wallet_helpers::new_address(*key_info.key_type(), &public_key)?;
34        Ok(Key {
35            key_info,
36            public_key,
37            address,
38        })
39    }
40}
41
42// This is a Wallet, it contains 2 HashMaps:
43// - keys which is a HashMap of Keys resolved by their Address
44// - keystore which is a HashMap of KeyInfos resolved by their Address
45/// A wallet is a collection of private keys with optional persistence and
46/// optional encryption.
47#[cfg(test)]
48#[derive(Clone, PartialEq, Debug, Eq)]
49pub struct Wallet {
50    keys: HashMap<Address, Key>,
51    keystore: KeyStore,
52}
53
54#[cfg(test)]
55impl Wallet {
56    /// Return a new wallet with a given `KeyStore`
57    pub fn new(keystore: KeyStore) -> Self {
58        Wallet {
59            keys: HashMap::new(),
60            keystore,
61        }
62    }
63
64    /// Return a wallet from a given amount of keys.
65    pub fn new_from_keys(keystore: KeyStore, key_vec: impl IntoIterator<Item = Key>) -> Self {
66        let mut keys: HashMap<Address, Key> = HashMap::new();
67        for item in key_vec.into_iter() {
68            keys.insert(item.address, item);
69        }
70        Wallet { keys, keystore }
71    }
72
73    // If this key does not exist in the keys hashmap, check if this key is in
74    // the keystore, if it is, then add it to keys, otherwise return Error
75    /// Return the key that is resolved by a given address,
76    pub fn find_key(&mut self, addr: &Address) -> Result<Key, Error> {
77        if let Some(k) = self.keys.get(addr) {
78            return Ok(k.clone());
79        }
80        let key = try_find_key(addr, &self.keystore)?;
81        self.keys.insert(*addr, key.clone());
82        Ok(key)
83    }
84
85    /// Return the resultant `Signature` after signing a given message
86    pub fn sign(&mut self, addr: &Address, msg: &[u8]) -> Result<Signature, Error> {
87        // this will return an error if the key cannot be found in either the keys
88        // hashmap or it is not found in the keystore
89        let key = self.find_key(addr).map_err(|_| Error::KeyNotExists)?;
90        wallet_helpers::sign(*key.key_info.key_type(), key.key_info.private_key(), msg)
91    }
92
93    /// Return the `KeyInfo` for a given address
94    pub fn export(&mut self, addr: &Address) -> Result<KeyInfo, Error> {
95        let k = self.find_key(addr)?;
96        Ok(k.key_info)
97    }
98
99    /// Add `KeyInfo` to the wallet, return the address that resolves to this
100    /// newly added `KeyInfo`
101    pub fn import(&mut self, key_info: KeyInfo) -> Result<Address, Error> {
102        let k = Key::try_from(key_info)?;
103        let addr = format!("wallet-{}", k.address);
104        self.keystore.put(&addr, k.key_info)?;
105        Ok(k.address)
106    }
107
108    /// Return a vector that contains all of the addresses in the wallet's
109    /// `KeyStore`
110    pub fn list_addrs(&self) -> Result<Vec<Address>, Error> {
111        list_addrs(&self.keystore)
112    }
113
114    /// Return the address of the default `KeyInfo` in the wallet
115    pub fn get_default(&self) -> Result<Address, Error> {
116        let key_info = self.keystore.get("default")?;
117        let k = Key::try_from(key_info)?;
118        Ok(k.address)
119    }
120
121    /// Set a default `KeyInfo` to the wallet
122    pub fn set_default(&mut self, addr: Address) -> anyhow::Result<()> {
123        let key_info = try_find(&addr, &self.keystore)?;
124        self.keystore.set_default(key_info)?;
125        Ok(())
126    }
127
128    /// Generate a new address that fits the requirement of the given
129    /// `SignatureType`
130    pub fn generate_addr(&mut self, typ: SignatureType) -> anyhow::Result<Address> {
131        let key = generate_key(typ)?;
132        let addr = format!("wallet-{}", key.address);
133        self.keystore.put(&addr, key.key_info.clone())?;
134        self.keys.insert(key.address, key.clone());
135        let value = self.keystore.get("default");
136        if value.is_err() {
137            self.keystore
138                .put("default", key.key_info.clone())
139                .map_err(|err| Error::Other(err.to_string()))?;
140        }
141
142        Ok(key.address)
143    }
144
145    /// Return whether or not the Wallet contains a key that is resolved by the
146    /// supplied address
147    pub fn has_key(&mut self, addr: &Address) -> bool {
148        self.find_key(addr).is_ok()
149    }
150}
151
152/// Return the default address for `KeyStore`
153pub fn get_default(keystore: &KeyStore) -> Result<Option<Address>, Error> {
154    if let Ok(key_info) = keystore.get("default") {
155        let k = Key::try_from(key_info)?;
156        Ok(Some(k.address))
157    } else {
158        Ok(None)
159    }
160}
161
162/// Return vector of addresses sorted by their string representation in
163/// `KeyStore`
164pub fn list_addrs(keystore: &KeyStore) -> Result<Vec<Address>, Error> {
165    let mut all = keystore.list();
166    all.sort();
167    let mut out = Vec::new();
168    for i in all {
169        if let Some(addr_str) = i.strip_prefix("wallet-")
170            && let Ok(addr) = crate::shim::address::StrictAddress::from_str(addr_str)
171        {
172            out.push(addr.into());
173        }
174    }
175    Ok(out)
176}
177
178/// Removes a key corresponding to given address
179pub fn remove_key(addr: &Address, keystore: &mut KeyStore) -> Result<(), Error> {
180    let key_string = format!("wallet-{addr}");
181    let deleted_keyinfo = keystore
182        .remove(&key_string)
183        .map_err(|_| Error::KeyNotExists)?;
184    if let Ok(default_keyinfo) = keystore.get("default")
185        && default_keyinfo == deleted_keyinfo
186    {
187        keystore
188            .remove("default")
189            .map_err(|_| Error::KeyNotExists)?;
190    }
191    println!("wallet {addr} deleted");
192    Ok(())
193}
194
195/// Returns key info corresponding to given address
196pub fn try_find(addr: &Address, keystore: &KeyStore) -> Result<KeyInfo, Error> {
197    let key_string = format!("wallet-{addr}");
198    match keystore.get(&key_string) {
199        Ok(k) => Ok(k),
200        Err(_) => {
201            let mut new_addr = addr.to_string();
202            if new_addr.len() < 2 {
203                return Err(Error::Other(format!("Invalid addr {new_addr}")));
204            }
205            // Try to replace prefix with testnet, for backwards compatibility
206            // * We might be able to remove this, look into variants
207            new_addr.replace_range(0..1, "t");
208            let key_string = format!("wallet-{new_addr}");
209            let key_info = match keystore.get(&key_string) {
210                Ok(k) => k,
211                Err(_) => keystore.get(&format!("wallet-f{}", &new_addr[1..]))?,
212            };
213            Ok(key_info)
214        }
215    }
216}
217
218pub fn try_find_key(addr: &Address, keystore: &KeyStore) -> Result<Key, Error> {
219    let ki = try_find(addr, keystore)?;
220    ki.try_into()
221}
222
223/// Return `KeyInfo` for given address in `KeyStore`
224pub fn export_key_info(addr: &Address, keystore: &KeyStore) -> Result<KeyInfo, Error> {
225    let key = try_find_key(addr, keystore)?;
226    Ok(key.key_info)
227}
228
229/// Generate new key of given `SignatureType`
230pub fn generate_key(typ: SignatureType) -> Result<Key, Error> {
231    let private_key = wallet_helpers::generate(typ)?;
232    let key_info = KeyInfo::new(typ, private_key);
233    Key::try_from(key_info)
234}
235
236#[cfg(test)]
237mod tests {
238    use crate::utils::encoding::{blake2b_256, keccak_256};
239    use bls_signatures::{PrivateKey as BlsPrivate, Serialize};
240
241    use super::*;
242    use crate::key_management::{KeyStoreConfig, generate};
243
244    fn construct_priv_keys() -> Vec<Key> {
245        let mut secp_keys = Vec::new();
246        let mut bls_keys = Vec::new();
247        let mut delegated_keys = Vec::new();
248        for _ in 1..5 {
249            let secp_priv_key = generate(SignatureType::Secp256k1).unwrap();
250            let secp_key_info = KeyInfo::new(SignatureType::Secp256k1, secp_priv_key);
251            let secp_key = Key::try_from(secp_key_info).unwrap();
252            secp_keys.push(secp_key);
253
254            let bls_priv_key = generate(SignatureType::Bls).unwrap();
255            let bls_key_info = KeyInfo::new(SignatureType::Bls, bls_priv_key);
256            let bls_key = Key::try_from(bls_key_info).unwrap();
257            bls_keys.push(bls_key);
258
259            let delegated_priv_key = generate(SignatureType::Delegated).unwrap();
260            let delegated_key_info = KeyInfo::new(SignatureType::Delegated, delegated_priv_key);
261            let delegated_key = Key::try_from(delegated_key_info).unwrap();
262            delegated_keys.push(delegated_key);
263        }
264
265        secp_keys.append(bls_keys.as_mut());
266        secp_keys.append(delegated_keys.as_mut());
267        secp_keys
268    }
269
270    fn generate_wallet() -> Wallet {
271        let key_vec = construct_priv_keys();
272        Wallet::new_from_keys(KeyStore::new(KeyStoreConfig::Memory).unwrap(), key_vec)
273    }
274
275    #[test]
276    fn contains_key() {
277        let key_vec = construct_priv_keys();
278        let found_key = key_vec[0].clone();
279        let addr = key_vec[0].address;
280
281        let mut wallet =
282            Wallet::new_from_keys(KeyStore::new(KeyStoreConfig::Memory).unwrap(), key_vec);
283
284        // make sure that this address resolves to the right key
285        assert_eq!(wallet.find_key(&addr).unwrap(), found_key);
286        // make sure that has_key returns true as well
287        assert!(wallet.has_key(&addr));
288
289        let new_priv_key = generate(SignatureType::Bls).unwrap();
290        let pub_key =
291            wallet_helpers::to_uncompressed_public_key(SignatureType::Bls, new_priv_key.as_slice())
292                .unwrap();
293        let address = Address::new_bls(pub_key.as_slice()).unwrap();
294
295        // test to see if the new key has been created and added to the wallet
296        assert!(!wallet.has_key(&address));
297        // test to make sure that the newly made key cannot be added to the wallet
298        // because it is not found in the keystore
299        assert!(matches!(
300            wallet.find_key(&address).unwrap_err(),
301            Error::KeyInfo
302        ));
303        // sanity check to make sure that the key has not been added to the wallet
304        assert!(!wallet.has_key(&address));
305    }
306
307    #[test]
308    fn secp_sign() {
309        let key_vec = construct_priv_keys();
310        let priv_key_bytes = key_vec[2].key_info.private_key().clone();
311        let addr = key_vec[2].address;
312
313        let keystore = KeyStore::new(KeyStoreConfig::Memory).unwrap();
314        let mut wallet = Wallet::new_from_keys(keystore, key_vec);
315        let msg = [0u8; 64];
316
317        let msg_sig = wallet.sign(&addr, &msg).unwrap();
318
319        let msg_complete = blake2b_256(&msg);
320        let priv_key = k256::ecdsa::SigningKey::from_slice(&priv_key_bytes).unwrap();
321        let (sig, recovery_id) = priv_key.sign_prehash_recoverable(&msg_complete).unwrap();
322        let mut new_bytes = [0; 65];
323        new_bytes[..64].copy_from_slice(&sig.to_bytes());
324        new_bytes[64] = recovery_id.to_byte();
325        let actual = Signature::new_secp256k1(new_bytes.to_vec());
326        assert_eq!(msg_sig, actual)
327    }
328
329    #[test]
330    fn bls_sign() {
331        let key_vec = construct_priv_keys();
332        let priv_key_bytes = key_vec[4].key_info.private_key().clone();
333        let addr = key_vec[4].address;
334        let mut wallet =
335            Wallet::new_from_keys(KeyStore::new(KeyStoreConfig::Memory).unwrap(), key_vec);
336
337        let msg = [0u8; 64];
338        let msg_sign = wallet.sign(&addr, &msg).unwrap();
339
340        let priv_key = BlsPrivate::from_bytes(&priv_key_bytes).unwrap();
341        let sig = priv_key.sign(msg);
342        let actual = Signature::new_bls(sig.as_bytes());
343        assert_eq!(msg_sign, actual);
344    }
345
346    #[test]
347    fn delegated_sign() {
348        let key_vec = construct_priv_keys();
349        let priv_key_bytes = key_vec[9].key_info.private_key().clone();
350        let addr = key_vec[9].address;
351
352        let keystore = KeyStore::new(KeyStoreConfig::Memory).unwrap();
353        let mut wallet = Wallet::new_from_keys(keystore, key_vec);
354        let msg = [0u8; 64];
355
356        let msg_sig = wallet.sign(&addr, &msg).unwrap();
357
358        let msg_complete = keccak_256(&msg);
359        let priv_key = k256::ecdsa::SigningKey::from_slice(&priv_key_bytes).unwrap();
360        let (sig, recovery_id) = priv_key.sign_prehash_recoverable(&msg_complete).unwrap();
361        let mut new_bytes = [0; 65];
362        new_bytes[..64].copy_from_slice(&sig.to_bytes());
363        new_bytes[64] = recovery_id.to_byte();
364        let actual = Signature::new_delegated(new_bytes.to_vec());
365        assert_eq!(msg_sig, actual)
366    }
367
368    #[test]
369    fn import_export() {
370        let key_vec = construct_priv_keys();
371        let key = key_vec[0].clone();
372        let keystore = KeyStore::new(KeyStoreConfig::Memory).unwrap();
373        let mut wallet = Wallet::new_from_keys(keystore, key_vec);
374
375        let key_info = wallet.export(&key.address).unwrap();
376        // test to see if export returns the correct key_info
377        assert_eq!(key_info, key.key_info);
378
379        let new_priv_key = generate(SignatureType::Secp256k1).unwrap();
380        let pub_key = wallet_helpers::to_uncompressed_public_key(
381            SignatureType::Secp256k1,
382            new_priv_key.as_slice(),
383        )
384        .unwrap();
385        let test_addr = Address::new_secp256k1(pub_key.as_slice()).unwrap();
386        let key_info_err = wallet.export(&test_addr).unwrap_err();
387        // test to make sure that an error is raised when an incorrect address is added
388        assert!(matches!(key_info_err, Error::KeyInfo));
389
390        let test_key_info = KeyInfo::new(SignatureType::Secp256k1, new_priv_key);
391        // make sure that key_info has been imported to wallet
392        assert!(wallet.import(test_key_info.clone()).is_ok());
393
394        let duplicate_error = wallet.import(test_key_info).unwrap_err();
395        // make sure that error is thrown when attempted to re-import a duplicate
396        // key_info
397        assert!(matches!(duplicate_error, Error::KeyExists));
398    }
399
400    #[test]
401    fn list_addr() {
402        let key_vec = construct_priv_keys();
403        let mut addr_string_vec = Vec::new();
404
405        let mut key_store = KeyStore::new(KeyStoreConfig::Memory).unwrap();
406
407        for i in &key_vec {
408            addr_string_vec.push(i.address.to_string());
409
410            let addr_string = format!("wallet-{}", i.address);
411            key_store.put(&addr_string, i.key_info.clone()).unwrap();
412        }
413
414        addr_string_vec.sort();
415
416        let mut addr_vec = Vec::new();
417
418        for addr in addr_string_vec {
419            addr_vec.push(Address::from_str(addr.as_str()).unwrap())
420        }
421
422        let wallet = Wallet::new(key_store);
423
424        let test_addr_vec = wallet.list_addrs().unwrap();
425
426        // check to see if the addrs in wallet are the same as the key_vec before it was
427        // added to the wallet
428        assert_eq!(test_addr_vec, addr_vec);
429    }
430
431    #[test]
432    fn generate_new_key() {
433        let mut wallet = generate_wallet();
434        let addr = wallet.generate_addr(SignatureType::Bls).unwrap();
435        let key = wallet.keystore.get("default").unwrap();
436        // make sure that the newly generated key is the default key - checking by key
437        // type
438        assert_eq!(&SignatureType::Bls, key.key_type());
439
440        let address = format!("wallet-{addr}");
441
442        let key_info = wallet.keystore.get(&address).unwrap();
443        let key = wallet.keys.get(&addr).unwrap();
444
445        // these assertions will make sure that the key has actually been added to the
446        // wallet
447        assert_eq!(key_info.key_type(), &SignatureType::Bls);
448        assert_eq!(key.address, addr);
449    }
450
451    #[test]
452    fn get_set_default() {
453        let key_store = KeyStore::new(KeyStoreConfig::Memory).unwrap();
454        let mut wallet = Wallet::new(key_store);
455        // check to make sure that there is no default
456        assert!(matches!(wallet.get_default().unwrap_err(), Error::KeyInfo));
457
458        let new_priv_key = generate(SignatureType::Secp256k1).unwrap();
459        let pub_key = wallet_helpers::to_uncompressed_public_key(
460            SignatureType::Secp256k1,
461            new_priv_key.as_slice(),
462        )
463        .unwrap();
464        let test_addr = Address::new_secp256k1(pub_key.as_slice()).unwrap();
465
466        let key_info = KeyInfo::new(SignatureType::Secp256k1, new_priv_key);
467        let test_addr_string = format!("wallet-{test_addr}");
468
469        wallet.keystore.put(&test_addr_string, key_info).unwrap();
470
471        // check to make sure that the set_default function completed without error
472        assert!(wallet.set_default(test_addr).is_ok());
473
474        // check to make sure that the test_addr is actually the default addr for the
475        // wallet
476        assert_eq!(wallet.get_default().unwrap(), test_addr);
477    }
478
479    #[test]
480    fn set_default_replaces_existing_default() {
481        let mut wallet = generate_wallet();
482        let addr_1 = wallet.generate_addr(SignatureType::Secp256k1).unwrap();
483        let addr_2 = wallet.generate_addr(SignatureType::Bls).unwrap();
484
485        // check to make sure that there is a default
486        assert_eq!(wallet.get_default().unwrap(), addr_1);
487        wallet.set_default(addr_2).unwrap();
488        // check to make sure that default is replaced
489        assert_eq!(wallet.get_default().unwrap(), addr_2);
490    }
491
492    #[test]
493    fn secp_verify() {
494        let secp_priv_key = generate(SignatureType::Secp256k1).unwrap();
495        let secp_key_info = KeyInfo::new(SignatureType::Secp256k1, secp_priv_key);
496        let secp_key = Key::try_from(secp_key_info).unwrap();
497        let addr = secp_key.address;
498        let key_store = KeyStore::new(KeyStoreConfig::Memory).unwrap();
499        let mut wallet = Wallet::new_from_keys(key_store, vec![secp_key]);
500
501        let msg = [0u8; 64];
502
503        let sig = wallet.sign(&addr, &msg).unwrap();
504        sig.verify(&msg, &addr).unwrap();
505
506        // invalid verify check
507        let invalid_addr = wallet.generate_addr(SignatureType::Secp256k1).unwrap();
508        assert!(sig.verify(&msg, &invalid_addr).is_err())
509    }
510
511    #[test]
512    fn bls_verify_test() {
513        let bls_priv_key = generate(SignatureType::Bls).unwrap();
514        let bls_key_info = KeyInfo::new(SignatureType::Bls, bls_priv_key);
515        let bls_key = Key::try_from(bls_key_info).unwrap();
516        let addr = bls_key.address;
517        let key_store = KeyStore::new(KeyStoreConfig::Memory).unwrap();
518        let mut wallet = Wallet::new_from_keys(key_store, vec![bls_key]);
519
520        let msg = [0u8; 64];
521
522        let sig = wallet.sign(&addr, &msg).unwrap();
523        sig.verify(&msg, &addr).unwrap();
524
525        // invalid verify check
526        let invalid_addr = wallet.generate_addr(SignatureType::Bls).unwrap();
527        assert!(sig.verify(&msg, &invalid_addr).is_err())
528    }
529
530    #[test]
531    fn delegated_verify() {
532        let delegated_priv_key = generate(SignatureType::Delegated).unwrap();
533        let delegated_key_info = KeyInfo::new(SignatureType::Delegated, delegated_priv_key);
534        let delegated_key = Key::try_from(delegated_key_info).unwrap();
535        let addr = delegated_key.address;
536
537        let key_store = KeyStore::new(KeyStoreConfig::Memory).unwrap();
538        let mut wallet = Wallet::new_from_keys(key_store, vec![delegated_key]);
539
540        let msg = [0u8; 64];
541
542        let sig = wallet.sign(&addr, &msg).unwrap();
543        sig.verify(&msg, &addr).unwrap();
544
545        // invalid verify check
546        let invalid_addr = wallet.generate_addr(SignatureType::Delegated).unwrap();
547        assert!(sig.verify(&msg, &invalid_addr).is_err())
548    }
549}