1#![warn(missing_docs)]
5#![warn(unused_crate_dependencies)]
6
7pub use error::WalletError;
8
9use massa_cipher::{decrypt, encrypt, CipherData, Salt};
10use massa_hash::Hash;
11use massa_models::address::Address;
12use massa_models::composite::PubkeySig;
13use massa_models::operation::{Operation, OperationSerializer, SecureShareOperation};
14use massa_models::prehash::{PreHashMap, PreHashSet};
15use massa_models::secure_share::SecureShareContent;
16use massa_signature::{KeyPair, PublicKey};
17use serde::{Deserialize, Serialize};
18use std::collections::hash_map::Entry;
19use std::collections::HashSet;
20use std::path::{Path, PathBuf};
21use std::str::FromStr;
22use zeroize::{Zeroize, ZeroizeOnDrop};
23
24mod error;
25
26const WALLET_VERSION: u64 = 1;
27
28#[derive(Clone, Debug, Deserialize, Serialize, Zeroize, ZeroizeOnDrop)]
30pub struct Wallet {
31 #[zeroize(skip)]
33 pub keys: PreHashMap<Address, KeyPair>,
34 #[zeroize(skip)]
36 wallet_path: PathBuf,
37 password: String,
39 chain_id: u64,
41}
42
43#[derive(Clone, Debug, Deserialize, Serialize)]
44#[serde(rename_all = "PascalCase")]
45struct WalletFileFormat {
47 version: u64,
48 nickname: String,
49 address: String,
50 salt: Salt,
51 nonce: [u8; 12],
52 ciphered_data: Vec<u8>,
53 public_key: Vec<u8>,
54}
55
56impl Wallet {
58 pub fn new(path: PathBuf, password: String, chain_id: u64) -> Result<Wallet, WalletError> {
60 if path.is_dir() {
61 let mut keys = PreHashMap::default();
62 for entry in std::fs::read_dir(&path)? {
63 let entry = entry?;
64 let path = entry.path();
65 if path.is_file() {
66 const WALLET_EXTENSIONS: [&str; 2] = ["yaml", "yml"];
67 let Some(ext) = path.extension().and_then(|e| e.to_str()) else {
68 continue;
69 };
70 if !WALLET_EXTENSIONS.contains(&ext.to_ascii_lowercase().as_str()) {
71 continue;
72 }
73 let content = &std::fs::read(&path)?[..];
74 let mut wallet = serde_yaml::from_slice::<WalletFileFormat>(content)?;
75 if wallet.version == 0 {
76 wallet.version = 1;
78 }
79 if wallet.version != WALLET_VERSION {
81 return Err(WalletError::VersionError(format!(
82 "Unsupported wallet version {}",
83 wallet.version
84 )));
85 }
86 let mut secret_key = decrypt(
87 &password,
88 CipherData {
89 salt: wallet.salt,
90 nonce: wallet.nonce,
91 encrypted_bytes: wallet.ciphered_data,
92 },
93 )?;
94 match secret_key.len() {
96 33 => {
97 },
99 65 => {
100 secret_key.truncate(33);
103 },
104 32 | 64 if wallet.version == 0 => {
105 return Err(WalletError::VersionError("Your wallet is from an old version that does not follow the standard. Please create a new wallet.".to_string()))
106 }
107 _ => {
108 return Err(WalletError::VersionError("Invalid wallet/version matching: your wallet does not follow its version's secret key encoding format.".to_string()))
109 }
110 }
111 let keypair = KeyPair::from_bytes(&secret_key)?;
112 let declared_address = Address::from_str(&wallet.address)?;
118 let derived_address = Address::from_public_key(&keypair.get_public_key());
119 if derived_address != declared_address {
120 return Err(WalletError::InconsistentWalletFile(format!(
121 "wallet file declares address {} but the decrypted key derives {}",
122 declared_address, derived_address
123 )));
124 }
125 keys.insert(derived_address, keypair);
126 }
127 }
128 Ok(Wallet {
129 keys,
130 wallet_path: path,
131 password,
132 chain_id,
133 })
134 } else {
135 let wallet = Wallet {
136 keys: PreHashMap::default(),
137 wallet_path: path,
138 password,
139 chain_id,
140 };
141 wallet.save()?;
142 Ok(wallet)
143 }
144 }
145
146 pub fn sign_message(&self, address: &Address, msg: Vec<u8>) -> Option<PubkeySig> {
150 if let Some(key) = self.keys.get(address) {
151 if let Ok(signature) = key.sign(&Hash::compute_from(&msg)) {
152 Some(PubkeySig {
153 public_key: key.get_public_key(),
154 signature,
155 })
156 } else {
157 None
158 }
159 } else {
160 None
161 }
162 }
163
164 pub fn add_keypairs(&mut self, keys: Vec<KeyPair>) -> Result<Vec<Address>, WalletError> {
167 let mut changed = false;
168 let mut addrs = Vec::with_capacity(keys.len());
169 for key in keys {
170 let addr = Address::from_public_key(&key.get_public_key());
171 if let Entry::Vacant(e) = self.keys.entry(addr) {
172 e.insert(key);
173 changed = true;
174 }
175 addrs.push(addr);
176 }
177 if changed {
178 self.save()?;
179 }
180 Ok(addrs)
181 }
182
183 pub fn remove_addresses(&mut self, addresses: &Vec<Address>) -> Result<bool, WalletError> {
186 let mut changed = false;
187 for address in addresses {
188 if self.keys.remove(address).is_some() {
189 changed = true;
190 }
191 }
192 Ok(changed)
193 }
194
195 pub fn find_associated_keypair(&self, address: &Address) -> Option<&KeyPair> {
197 self.keys.get(address)
198 }
199
200 pub fn find_associated_public_key(&self, address: &Address) -> Option<PublicKey> {
202 self.keys
203 .get(address)
204 .map(|keypair| keypair.get_public_key())
205 }
206
207 pub fn get_wallet_address_list(&self) -> PreHashSet<Address> {
209 self.keys.keys().copied().collect()
210 }
211
212 fn is_managed_wallet_file(path: &Path) -> bool {
219 if !path.is_file() {
220 return false;
221 }
222 let ext_ok = path
223 .extension()
224 .and_then(|e| e.to_str())
225 .map(|e| {
226 let e = e.to_ascii_lowercase();
227 e == "yaml" || e == "yml"
228 })
229 .unwrap_or(false);
230 let name_ok = path
231 .file_name()
232 .and_then(|n| n.to_str())
233 .map(|n| n.starts_with("wallet_"))
234 .unwrap_or(false);
235 ext_ok && name_ok
236 }
237
238 pub fn save(&self) -> Result<(), WalletError> {
240 let mut existing_keys: HashSet<PathBuf> = HashSet::new();
241 if !self.wallet_path.exists() {
242 std::fs::create_dir_all(&self.wallet_path)?;
243 } else {
244 let read_dir = std::fs::read_dir(&self.wallet_path)?;
245 for path in read_dir {
246 let path = path?.path();
247 if Self::is_managed_wallet_file(&path) {
250 existing_keys.insert(path);
251 }
252 }
253 }
254 let mut persisted_keys: HashSet<PathBuf> = HashSet::new();
255 for (addr, keypair) in &self.keys {
257 let encrypted_secret = encrypt(&self.password, &keypair.to_bytes())?;
258 let file_formatted = WalletFileFormat {
259 version: WALLET_VERSION,
260 nickname: addr.to_string(),
261 address: addr.to_string(),
262 salt: encrypted_secret.salt,
263 nonce: encrypted_secret.nonce,
264 ciphered_data: encrypted_secret.encrypted_bytes,
265 public_key: keypair.get_public_key().to_bytes().to_vec(),
266 };
267 let ser_keys = serde_yaml::to_string(&file_formatted)?;
268 let file_path = self.wallet_path.join(format!("wallet_{}.yaml", addr));
269
270 std::fs::write(&file_path, ser_keys)?;
271 persisted_keys.insert(file_path);
272 }
273
274 let to_remove = existing_keys.difference(&persisted_keys);
275 for path in to_remove {
276 std::fs::remove_file(path)?;
277 }
278
279 Ok(())
280 }
281
282 pub fn get_full_wallet(&self) -> &PreHashMap<Address, KeyPair> {
284 &self.keys
285 }
286
287 pub fn create_operation(
289 &self,
290 content: Operation,
291 address: Address,
292 ) -> Result<SecureShareOperation, WalletError> {
293 let sender_keypair = self
294 .find_associated_keypair(&address)
295 .ok_or(WalletError::MissingKeyError(address))?;
296 Ok(Operation::new_verifiable(
297 content,
298 OperationSerializer::new(),
299 sender_keypair,
300 self.chain_id,
301 )
302 .unwrap())
303 }
304}
305
306impl std::fmt::Display for Wallet {
307 fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
308 writeln!(f)?;
309 for (addr, keypair) in &self.keys {
310 writeln!(f, "Secret key: {}", keypair)?;
311 writeln!(f, "Public key: {}", keypair.get_public_key())?;
312 writeln!(f, "Address: {}", addr)?;
313 }
314 Ok(())
315 }
316}
317
318#[cfg(feature = "test-exports")]
320pub mod test_exports;
321
322#[cfg(all(test, feature = "test-exports"))]
323mod tests {
324 use super::*;
325 use massa_signature::KeyPair;
326 use tempfile::TempDir;
327
328 #[test]
329 fn save_only_removes_managed_wallet_files() {
330 let dir = TempDir::new().unwrap();
331 let dir_path = dir.path().to_path_buf();
332
333 let notes = dir_path.join("notes.txt");
335 let backup = dir_path.join("backup.yaml"); let stale = dir_path.join("wallet_stale.yaml"); std::fs::write(¬es, b"important recovery notes").unwrap();
338 std::fs::write(&backup, b"not: a-wallet").unwrap();
339 std::fs::write(&stale, b"stale: wallet").unwrap();
340
341 let wallet = Wallet {
342 keys: PreHashMap::default(),
343 wallet_path: dir_path.clone(),
344 password: "pw".to_string(),
345 chain_id: 0,
346 };
347 wallet.save().unwrap();
348
349 assert!(notes.exists(), "unrelated non-yaml file must be preserved");
350 assert!(backup.exists(), "non-wallet .yaml file must be preserved");
351 assert!(
352 !stale.exists(),
353 "stale managed wallet_*.yaml file must be removed"
354 );
355 }
356
357 #[test]
358 fn honest_wallet_loads_and_tampered_address_is_rejected() {
359 let dir = TempDir::new().unwrap();
360 let dir_path = dir.path().to_path_buf();
361
362 let mut wallet = Wallet::new(dir_path.clone(), "pw".to_string(), 0).unwrap();
364 let kp = KeyPair::generate(0).unwrap();
365 let addr = Address::from_public_key(&kp.get_public_key());
366 wallet.add_keypairs(vec![kp]).unwrap();
367
368 let reloaded = Wallet::new(dir_path.clone(), "pw".to_string(), 0).unwrap();
370 assert!(reloaded.keys.contains_key(&addr));
371
372 let file = dir_path.join(format!("wallet_{}.yaml", addr));
375 let content = std::fs::read(&file).unwrap();
376 let mut parsed: WalletFileFormat = serde_yaml::from_slice(&content).unwrap();
377 let other = Address::from_public_key(&KeyPair::generate(0).unwrap().get_public_key());
378 parsed.address = other.to_string();
379 std::fs::write(&file, serde_yaml::to_string(&parsed).unwrap()).unwrap();
380
381 let err = Wallet::new(dir_path, "pw".to_string(), 0)
384 .expect_err("a tampered declared address must be rejected");
385 assert!(matches!(err, WalletError::InconsistentWalletFile(_)));
386 }
387}