Skip to main content

origin_crypto_sdk/
stealth.rs

1// SPDX-License-Identifier: Apache-2.0
2
3//! Stealth address primitives.
4//!
5//! # Submodules
6//!
7//! - [`kdf`] — Key derivation (master keys + per-index subkeys)
8//! - [`stealth::pow`] — Hashcash-based proof-of-work to rate-limit enumeration
9
10pub mod kdf {
11    use crate::error::{CryptoError, Result};
12    use crate::kdf::hkdf::hkdf_sha3_256;
13    use crate::seed::SeedHandle;
14
15    /// Master keys for stealth address derivation.
16    #[derive(Debug, Clone)]
17    pub struct StealthMasterKeys {
18        /// Master viewing key (for deriving per-address viewing secrets).
19        pub viewing: [u8; 32],
20        /// Master spending key (for deriving per-address spending secrets).
21        pub spending: [u8; 32],
22        /// Master ephemeral key (for deriving per-address ephemeral secrets).
23        pub ephemeral: [u8; 32],
24    }
25
26    /// Per-address stealth keys derived at a specific index.
27    #[derive(Debug, Clone)]
28    pub struct StealthAddressKeys {
29        /// Viewing secret key for this address.
30        pub viewing_secret: [u8; 32],
31        /// Spending secret key for this address.
32        pub spending_secret: [u8; 32],
33        /// Ephemeral secret key for this address.
34        pub ephemeral_secret: [u8; 32],
35    }
36
37    /// Derive stealth master keys from a seed handle.
38    ///
39    /// Uses three separate HKDF derivations with domain separation:
40    /// - Viewing:   HKDF(seed, salt="stealth:viewing:master",   info="origin-stealth-v1")
41    /// - Spending:  HKDF(seed, salt="stealth:spending:master",  info="origin-stealth-v1")
42    /// - Ephemeral: HKDF(seed, salt="stealth:ephemeral:master", info="origin-stealth-v1")
43    pub fn derive_stealth_master(seed: &SeedHandle) -> Result<StealthMasterKeys> {
44        let seed_bytes = seed
45            .as_bytes()
46            .ok_or_else(|| CryptoError::InvalidParameter("Seed handle expired".into()))?;
47
48        let mut viewing = [0u8; 32];
49        let mut spending = [0u8; 32];
50        let mut ephemeral = [0u8; 32];
51
52        hkdf_sha3_256(
53            seed_bytes,
54            Some(b"stealth:viewing:master"),
55            b"origin-stealth-v1",
56            &mut viewing,
57        )?;
58        hkdf_sha3_256(
59            seed_bytes,
60            Some(b"stealth:spending:master"),
61            b"origin-stealth-v1",
62            &mut spending,
63        )?;
64        hkdf_sha3_256(
65            seed_bytes,
66            Some(b"stealth:ephemeral:master"),
67            b"origin-stealth-v1",
68            &mut ephemeral,
69        )?;
70
71        Ok(StealthMasterKeys {
72            viewing,
73            spending,
74            ephemeral,
75        })
76    }
77
78    /// Derive stealth address keys at a specific index from master keys.
79    ///
80    /// Each key is derived as:
81    /// HKDF(master_key, salt=index_be_bytes || context, info="stealth-derive-v1")
82    pub fn derive_stealth_at_index(
83        master: &StealthMasterKeys,
84        index: u64,
85    ) -> Result<StealthAddressKeys> {
86        let index_bytes = index.to_be_bytes();
87
88        let mut viewing_secret = [0u8; 32];
89        let mut spending_secret = [0u8; 32];
90        let mut ephemeral_secret = [0u8; 32];
91
92        // Viewing: salt = index || "viewing"
93        let mut salt_v = [0u8; 40];
94        salt_v[0..8].copy_from_slice(&index_bytes);
95        salt_v[8..15].copy_from_slice(b"viewing");
96        hkdf_sha3_256(
97            &master.viewing,
98            Some(&salt_v),
99            b"stealth-derive-v1",
100            &mut viewing_secret,
101        )?;
102
103        // Spending: salt = index || "spending"
104        let mut salt_s = [0u8; 40];
105        salt_s[0..8].copy_from_slice(&index_bytes);
106        salt_s[8..16].copy_from_slice(b"spending");
107        hkdf_sha3_256(
108            &master.spending,
109            Some(&salt_s),
110            b"stealth-derive-v1",
111            &mut spending_secret,
112        )?;
113
114        // Ephemeral: salt = index || "ephemeral"
115        let mut salt_e = [0u8; 40];
116        salt_e[0..8].copy_from_slice(&index_bytes);
117        salt_e[8..17].copy_from_slice(b"ephemeral");
118        hkdf_sha3_256(
119            &master.ephemeral,
120            Some(&salt_e),
121            b"stealth-derive-v1",
122            &mut ephemeral_secret,
123        )?;
124
125        Ok(StealthAddressKeys {
126            viewing_secret,
127            spending_secret,
128            ephemeral_secret,
129        })
130    }
131
132    /// Derive stealth keys directly from a seed at a specific index.
133    ///
134    /// Convenience function that combines `derive_stealth_master` and
135    /// `derive_stealth_at_index`.
136    pub fn derive_stealth_from_seed(seed: &SeedHandle, index: u64) -> Result<StealthAddressKeys> {
137        let master = derive_stealth_master(seed)?;
138        derive_stealth_at_index(&master, index)
139    }
140
141    #[cfg(test)]
142    mod tests {
143        use super::*;
144        use std::time::Duration;
145
146        #[test]
147        fn test_derive_stealth_master_deterministic() {
148            let seed = SeedHandle::new(&[42u8; 32], None);
149            let master1 = derive_stealth_master(&seed).unwrap();
150            let master2 = derive_stealth_master(&seed).unwrap();
151
152            assert_eq!(master1.viewing, master2.viewing);
153            assert_eq!(master1.spending, master2.spending);
154            assert_eq!(master1.ephemeral, master2.ephemeral);
155        }
156
157        #[test]
158        fn test_different_seeds_different_masters() {
159            let seed1 = SeedHandle::new(&[1u8; 32], None);
160            let seed2 = SeedHandle::new(&[2u8; 32], None);
161
162            let master1 = derive_stealth_master(&seed1).unwrap();
163            let master2 = derive_stealth_master(&seed2).unwrap();
164
165            assert_ne!(master1.viewing, master2.viewing);
166            assert_ne!(master1.spending, master2.spending);
167            assert_ne!(master1.ephemeral, master2.ephemeral);
168        }
169
170        #[test]
171        fn test_different_indices_different_keys() {
172            let seed = SeedHandle::new(&[42u8; 32], None);
173            let master = derive_stealth_master(&seed).unwrap();
174
175            let keys0 = derive_stealth_at_index(&master, 0).unwrap();
176            let keys1 = derive_stealth_at_index(&master, 1).unwrap();
177            let keys2 = derive_stealth_at_index(&master, 2).unwrap();
178
179            assert_ne!(keys0.viewing_secret, keys1.viewing_secret);
180            assert_ne!(keys1.viewing_secret, keys2.viewing_secret);
181            assert_ne!(keys0.spending_secret, keys1.spending_secret);
182            assert_ne!(keys0.ephemeral_secret, keys1.ephemeral_secret);
183        }
184
185        #[test]
186        fn test_derive_at_index_deterministic() {
187            let seed = SeedHandle::new(&[42u8; 32], None);
188            let master = derive_stealth_master(&seed).unwrap();
189
190            let keys_a = derive_stealth_at_index(&master, 5).unwrap();
191            let keys_b = derive_stealth_at_index(&master, 5).unwrap();
192
193            assert_eq!(keys_a.viewing_secret, keys_b.viewing_secret);
194            assert_eq!(keys_a.spending_secret, keys_b.spending_secret);
195            assert_eq!(keys_a.ephemeral_secret, keys_b.ephemeral_secret);
196        }
197
198        #[test]
199        fn test_derive_from_seed_convenience() {
200            let seed = SeedHandle::new(&[42u8; 32], None);
201
202            let direct = derive_stealth_from_seed(&seed, 3).unwrap();
203            let master = derive_stealth_master(&seed).unwrap();
204            let stepped = derive_stealth_at_index(&master, 3).unwrap();
205
206            assert_eq!(direct.viewing_secret, stepped.viewing_secret);
207            assert_eq!(direct.spending_secret, stepped.spending_secret);
208            assert_eq!(direct.ephemeral_secret, stepped.ephemeral_secret);
209        }
210
211        #[test]
212        fn test_expired_seed_fails() {
213            let seed = SeedHandle::new(&[42u8; 32], Some(Duration::from_nanos(1)));
214            // Let it expire
215            std::thread::sleep(Duration::from_millis(10));
216            assert!(derive_stealth_master(&seed).is_err());
217        }
218
219        #[test]
220        fn test_master_keys_are_different_from_each_other() {
221            let seed = SeedHandle::new(&[42u8; 32], None);
222            let master = derive_stealth_master(&seed).unwrap();
223
224            // The three master keys should all be different
225            assert_ne!(master.viewing, master.spending);
226            assert_ne!(master.viewing, master.ephemeral);
227            assert_ne!(master.spending, master.ephemeral);
228        }
229
230        #[test]
231        fn test_large_index() {
232            let seed = SeedHandle::new(&[42u8; 32], None);
233            let master = derive_stealth_master(&seed).unwrap();
234
235            let keys = derive_stealth_at_index(&master, u64::MAX).unwrap();
236            // Should not panic and should produce valid-looking keys
237            assert!(keys.viewing_secret.iter().any(|&b| b != 0));
238        }
239    }
240}
241
242pub mod pow {
243    use crate::error::{CryptoError, Result};
244    use crate::primitives::sha3::sha3_256;
245
246    /// Configuration for the stealth PoW system.
247    #[derive(Debug, Clone)]
248    pub struct StealthPowConfig {
249        /// Base difficulty (leading zero bits required). Default: 20.
250        pub base_difficulty: u32,
251        /// Difficulty increment per address generated. Default: 0.
252        pub per_address_increment: u32,
253        /// Maximum difficulty cap. Default: 32.
254        pub max_difficulty: u32,
255    }
256
257    impl Default for StealthPowConfig {
258        fn default() -> Self {
259            Self {
260                base_difficulty: 20,
261                per_address_increment: 0,
262                max_difficulty: 32,
263            }
264        }
265    }
266
267    /// A proof-of-work solution for stealth address generation.
268    #[derive(Debug, Clone)]
269    pub struct StealthPowProof {
270        /// The nonce that satisfies the difficulty requirement.
271        pub nonce: [u8; 32],
272        /// Extra data (can be used for additional binding).
273        pub extra: [u8; 16],
274        /// Counter incremented during mining.
275        pub counter: u64,
276        /// The difficulty this proof was mined at.
277        pub difficulty: u32,
278    }
279
280    /// Solve the Hashcash PoW for stealth address generation.
281    ///
282    /// # Arguments
283    /// * `identity_pk`     — The identity's public key (for binding)
284    /// * `destination_hint` — A hint about the destination (e.g., recipient's pubkey hash)
285    /// * `difficulty`      — Number of leading zero bits required
286    ///
287    /// # Returns
288    /// A `StealthPowProof` and the number of hash iterations performed.
289    pub fn solve(
290        identity_pk: &[u8],
291        destination_hint: &[u8],
292        difficulty: u32,
293    ) -> Result<(StealthPowProof, u64)> {
294        if difficulty == 0 {
295            // Trivial: any nonce works
296            let nonce = [0u8; 32];
297            let extra = [0u8; 16];
298            return Ok((
299                StealthPowProof {
300                    nonce,
301                    extra,
302                    counter: 0,
303                    difficulty: 0,
304                },
305                0,
306            ));
307        }
308
309        if difficulty > 32 {
310            return Err(CryptoError::InvalidParameter(
311                "Difficulty cannot exceed 32 bits".into(),
312            ));
313        }
314
315        let target = compute_target(difficulty);
316        let mut counter: u64 = 0;
317        let mut nonce = [0u8; 32];
318        let mut extra = [0u8; 16];
319
320        // Use a deterministic but varied starting nonce
321        let mut seed_input = Vec::with_capacity(16 + identity_pk.len() + destination_hint.len());
322        seed_input.extend_from_slice(b"stealth-pow-seed");
323        seed_input.extend_from_slice(identity_pk);
324        seed_input.extend_from_slice(destination_hint);
325        let seed_hash = sha3_256(&seed_input);
326        nonce.copy_from_slice(&seed_hash);
327        extra.copy_from_slice(&seed_hash[0..16]);
328
329        loop {
330            let hash = compute_hash(identity_pk, destination_hint, &nonce, &extra, counter);
331
332            if meets_target(&hash, &target, difficulty) {
333                return Ok((
334                    StealthPowProof {
335                        nonce,
336                        extra,
337                        counter,
338                        difficulty,
339                    },
340                    counter + 1,
341                ));
342            }
343
344            counter += 1;
345
346            // Increment nonce as a big-endian counter (wraps around)
347            for i in (0..32).rev() {
348                if nonce[i] == 0xff {
349                    nonce[i] = 0;
350                } else {
351                    nonce[i] += 1;
352                    break;
353                }
354            }
355
356            // Safety: prevent infinite loops in test environments
357            if counter > 100_000_000 {
358                return Err(CryptoError::InvalidParameter(
359                    "PoW solve exceeded maximum iterations".into(),
360                ));
361            }
362        }
363    }
364
365    /// Verify a stealth PoW proof.
366    ///
367    /// # Arguments
368    /// * `proof`           — The proof to verify
369    /// * `identity_pk`     — The identity's public key
370    /// * `destination_hint` — The destination hint used during solving
371    ///
372    /// # Returns
373    /// `Ok(true)` if the proof is valid at the claimed difficulty.
374    pub fn verify(
375        proof: &StealthPowProof,
376        identity_pk: &[u8],
377        destination_hint: &[u8],
378    ) -> Result<bool> {
379        if proof.difficulty == 0 {
380            return Ok(true);
381        }
382
383        if proof.difficulty > 32 {
384            return Ok(false);
385        }
386
387        let target = compute_target(proof.difficulty);
388        let hash = compute_hash(
389            identity_pk,
390            destination_hint,
391            &proof.nonce,
392            &proof.extra,
393            proof.counter,
394        );
395
396        Ok(meets_target(&hash, &target, proof.difficulty))
397    }
398
399    /// Compute the effective difficulty for a given address index.
400    pub fn effective_difficulty(config: &StealthPowConfig, address_index: u64) -> u32 {
401        let effective =
402            config.base_difficulty + config.per_address_increment * address_index as u32;
403        effective.min(config.max_difficulty)
404    }
405
406    // --- Internal functions ---
407
408    fn compute_hash(
409        identity_pk: &[u8],
410        destination_hint: &[u8],
411        nonce: &[u8; 32],
412        extra: &[u8; 16],
413        counter: u64,
414    ) -> [u8; 32] {
415        let mut input =
416            Vec::with_capacity(14 + identity_pk.len() + destination_hint.len() + 32 + 16 + 8);
417        input.extend_from_slice(b"stealth-pow-v1");
418        input.extend_from_slice(identity_pk);
419        input.extend_from_slice(destination_hint);
420        input.extend_from_slice(nonce);
421        input.extend_from_slice(extra);
422        input.extend_from_slice(&counter.to_be_bytes());
423        sha3_256(&input)
424    }
425
426    fn compute_target(difficulty: u32) -> [u8; 32] {
427        let mut target = [0xffu8; 32];
428        let full_bytes = (difficulty / 8) as usize;
429        let remaining_bits = difficulty % 8;
430
431        for i in 0..full_bytes {
432            target[i] = 0x00;
433        }
434
435        if full_bytes < 32 && remaining_bits > 0 {
436            let mask = 0xffu8 >> remaining_bits;
437            target[full_bytes] = mask;
438        }
439
440        target
441    }
442
443    fn meets_target(hash: &[u8; 32], _target: &[u8; 32], difficulty: u32) -> bool {
444        let full_bytes = (difficulty / 8) as usize;
445        let remaining_bits = difficulty % 8;
446
447        // Check full zero bytes
448        for i in 0..full_bytes {
449            if hash[i] != 0x00 {
450                return false;
451            }
452        }
453
454        // Check remaining bits
455        if full_bytes < 32 && remaining_bits > 0 {
456            let mask = 0xffu8 >> remaining_bits;
457            if hash[full_bytes] & mask != 0 {
458                return false;
459            }
460        }
461
462        true
463    }
464
465    #[cfg(test)]
466    mod tests {
467        use super::*;
468
469        #[test]
470        fn test_solve_verify_roundtrip() {
471            let pk = [42u8; 32];
472            let hint = b"destination hint";
473            let difficulty = 16; // Low difficulty for fast tests
474
475            let (proof, iterations) = solve(&pk, hint, difficulty).unwrap();
476            assert!(iterations > 0 || difficulty == 0);
477            assert_eq!(proof.difficulty, difficulty);
478
479            assert!(verify(&proof, &pk, hint).unwrap());
480        }
481
482        #[test]
483        fn test_verify_wrong_pk_fails() {
484            let pk = [42u8; 32];
485            let wrong_pk = [99u8; 32];
486            let hint = b"hint";
487            let difficulty = 16;
488
489            let (proof, _) = solve(&pk, hint, difficulty).unwrap();
490            assert!(!verify(&proof, &wrong_pk, hint).unwrap());
491        }
492
493        #[test]
494        fn test_verify_wrong_hint_fails() {
495            let pk = [42u8; 32];
496            let hint = b"correct hint";
497            let wrong_hint = b"wrong hint";
498            let difficulty = 16;
499
500            let (proof, _) = solve(&pk, hint, difficulty).unwrap();
501            assert!(!verify(&proof, &pk, wrong_hint).unwrap());
502        }
503
504        #[test]
505        fn test_zero_difficulty_always_valid() {
506            let pk = [42u8; 32];
507            let hint = b"hint";
508
509            let (proof, iterations) = solve(&pk, hint, 0).unwrap();
510            assert_eq!(iterations, 0);
511            assert!(verify(&proof, &pk, hint).unwrap());
512        }
513
514        #[test]
515        fn test_difficulty_scaling() {
516            let pk = [42u8; 32];
517            let hint = b"hint";
518
519            // Higher difficulty should take more iterations
520            let (_, iters_12) = solve(&pk, hint, 12).unwrap();
521            let (_, iters_16) = solve(&pk, hint, 16).unwrap();
522            // Not guaranteed but very likely
523            assert!(
524                iters_16 >= iters_12,
525                "Higher difficulty should take >= iterations"
526            );
527        }
528
529        #[test]
530        fn test_effective_difficulty() {
531            let config = StealthPowConfig {
532                base_difficulty: 20,
533                per_address_increment: 1,
534                max_difficulty: 25,
535            };
536
537            assert_eq!(effective_difficulty(&config, 0), 20);
538            assert_eq!(effective_difficulty(&config, 3), 23);
539            assert_eq!(effective_difficulty(&config, 10), 25); // Capped at max
540        }
541
542        #[test]
543        fn test_different_proofs_for_different_inputs() {
544            let pk1 = [1u8; 32];
545            let pk2 = [2u8; 32];
546            let hint = b"same hint";
547            let difficulty = 16;
548
549            let (proof1, _) = solve(&pk1, hint, difficulty).unwrap();
550            let (proof2, _) = solve(&pk2, hint, difficulty).unwrap();
551
552            // Nonces should be different (derived from different seeds)
553            assert_ne!(proof1.nonce, proof2.nonce);
554        }
555    }
556}