1use crate::error::Error;
2use bitcoin::bip32::DerivationPath;
3use bitcoin::bip32::Xpriv;
4use bitcoin::key::Keypair;
5use bitcoin::secp256k1::Secp256k1;
6use std::sync::Arc;
7
8pub enum KeypairIndex {
9 New,
11 LastUnused,
13}
14
15pub trait KeyProvider: Send + Sync {
23 fn get_next_keypair(&self, keypair_index: KeypairIndex) -> Result<Keypair, Error>;
38
39 fn get_keypair_for_path(&self, path: &[u32]) -> Result<Keypair, Error>;
49
50 fn get_keypair_for_pk(&self, pk: &bitcoin::XOnlyPublicKey) -> Result<Keypair, Error>;
63
64 fn get_cached_pks(&self) -> Result<Vec<bitcoin::XOnlyPublicKey>, Error>;
77}
78
79pub trait DiscoverableKeyProvider: KeyProvider {
81 fn get_derivation_index_for_pk(&self, pk: &bitcoin::XOnlyPublicKey) -> Option<u32>;
83
84 fn derive_at_discovery_index(&self, index: u32) -> Result<Option<Keypair>, Error>;
86
87 fn cache_discovered_keypair(&self, index: u32, kp: Keypair) -> Result<(), Error>;
92
93 fn cache_keypair_at_index(&self, index: u32) -> Result<(), Error>;
99
100 fn mark_as_used(&self, pk: &bitcoin::XOnlyPublicKey) -> Result<(), Error>;
101}
102
103#[derive(Clone)]
108pub struct StaticKeyProvider {
109 kp: Keypair,
110}
111
112impl StaticKeyProvider {
113 pub fn new(kp: Keypair) -> Self {
115 Self { kp }
116 }
117}
118
119impl KeyProvider for StaticKeyProvider {
120 fn get_next_keypair(&self, _: KeypairIndex) -> Result<Keypair, Error> {
121 Ok(self.kp)
123 }
124
125 fn get_keypair_for_path(&self, _path: &[u32]) -> Result<Keypair, Error> {
126 Ok(self.kp)
128 }
129
130 fn get_keypair_for_pk(&self, pk: &bitcoin::XOnlyPublicKey) -> Result<Keypair, Error> {
131 let our_pk = self.kp.x_only_public_key().0;
133 if &our_pk == pk {
134 Ok(self.kp)
135 } else {
136 Err(Error::ad_hoc(format!(
137 "Public key mismatch: requested {pk}, but only have {our_pk}"
138 )))
139 }
140 }
141
142 fn get_cached_pks(&self) -> Result<Vec<bitcoin::XOnlyPublicKey>, Error> {
143 Ok(vec![self.kp.public_key().into()])
144 }
145}
146
147pub struct Bip32KeyProvider {
182 master_key: Xpriv,
183 base_path: DerivationPath,
184 next_index: Arc<std::sync::Mutex<u32>>,
186 key_cache:
189 Arc<std::sync::RwLock<std::collections::HashMap<bitcoin::XOnlyPublicKey, KeyCacheValue>>>,
190}
191
192#[derive(Clone, Copy)]
193pub struct KeyCacheValue {
194 path_index: u32,
195 kp: Keypair,
196 used: bool,
198}
199
200impl Bip32KeyProvider {
201 pub fn new(master_key: Xpriv, base_path: DerivationPath) -> Self {
209 Self {
210 master_key,
211 base_path,
212 next_index: Arc::new(std::sync::Mutex::new(0)),
213 key_cache: Arc::new(std::sync::RwLock::new(std::collections::HashMap::new())),
214 }
215 }
216
217 pub fn new_with_index(master_key: Xpriv, base_path: DerivationPath, start_index: u32) -> Self {
225 Self {
226 master_key,
227 base_path,
228 next_index: Arc::new(std::sync::Mutex::new(start_index)),
229 key_cache: Arc::new(std::sync::RwLock::new(std::collections::HashMap::new())),
230 }
231 }
232
233 fn derive_keypair(&self, path: &DerivationPath) -> Result<Keypair, Error> {
235 let secp = Secp256k1::new();
236 let derived_key = self
237 .master_key
238 .derive_priv(&secp, path)
239 .map_err(|e| Error::ad_hoc(format!("BIP32 derivation failed: {e}")))?;
240
241 Ok(derived_key.to_keypair(&secp))
242 }
243
244 fn derive_at_index(&self, index: u32) -> Result<Keypair, Error> {
246 use bitcoin::bip32::ChildNumber;
247
248 let path = self.base_path.clone();
249 let path = path.extend([ChildNumber::Normal { index }]);
250
251 self.derive_keypair(&path)
252 }
253}
254
255impl KeyProvider for Bip32KeyProvider {
256 fn get_next_keypair(&self, keypair_index: KeypairIndex) -> Result<Keypair, Error> {
257 match keypair_index {
258 KeypairIndex::New => {
259 let index = {
261 let mut next_index = self
262 .next_index
263 .lock()
264 .map_err(|e| Error::ad_hoc(format!("Failed to lock next_index: {e}")))?;
265 let current = *next_index;
266 *next_index = next_index
267 .checked_add(1)
268 .ok_or_else(|| Error::ad_hoc("Key derivation index overflow"))?;
269 current
270 };
271
272 let kp = self.derive_at_index(index)?;
274
275 let pk = kp.x_only_public_key().0;
277 {
278 let mut cache = self
279 .key_cache
280 .write()
281 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
282 cache.insert(
283 pk,
284 KeyCacheValue {
285 path_index: index,
286 kp,
287 used: false,
288 },
289 );
290 }
291
292 Ok(kp)
293 }
294 KeypairIndex::LastUnused => {
295 {
297 let cache = self
298 .key_cache
299 .read()
300 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
301
302 let unused = cache
304 .values()
305 .filter(|KeyCacheValue { used, .. }| !used)
306 .min_by_key(|KeyCacheValue { path_index, .. }| *path_index);
307
308 if let Some(KeyCacheValue { kp, .. }) = unused {
309 return Ok(*kp);
310 }
311 }
312
313 self.get_next_keypair(KeypairIndex::New)
315 }
316 }
317 }
318
319 fn get_keypair_for_path(&self, path: &[u32]) -> Result<Keypair, Error> {
320 use bitcoin::bip32::ChildNumber;
321 let child_numbers: Vec<ChildNumber> = path
322 .iter()
323 .map(|&n| {
324 if n & 0x8000_0000 != 0 {
325 ChildNumber::Hardened {
326 index: n & 0x7FFF_FFFF,
327 }
328 } else {
329 ChildNumber::Normal { index: n }
330 }
331 })
332 .collect();
333 let derivation_path = DerivationPath::from(child_numbers);
334 self.derive_keypair(&derivation_path)
335 }
336
337 fn get_keypair_for_pk(&self, pk: &bitcoin::XOnlyPublicKey) -> Result<Keypair, Error> {
338 {
340 let cache = self
341 .key_cache
342 .read()
343 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
344 if let Some(KeyCacheValue { kp, .. }) = cache.get(pk) {
345 return Ok(*kp);
346 }
347 }
348
349 let current_index = {
351 let next_index = self
352 .next_index
353 .lock()
354 .map_err(|e| Error::ad_hoc(format!("Failed to lock next_index: {e}")))?;
355 *next_index
356 };
357
358 for i in 0..current_index {
360 let kp = self.derive_at_index(i)?;
361 let derived_pk = kp.x_only_public_key().0;
362
363 if &derived_pk == pk {
364 let mut cache = self
366 .key_cache
367 .write()
368 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
369 cache.insert(
370 derived_pk,
371 KeyCacheValue {
372 path_index: i,
373 kp,
374 used: true,
375 },
376 );
377 return Ok(kp);
378 }
379 }
380
381 Err(Error::ad_hoc(format!(
382 "Public key {pk} not found in HD wallet. \
383 Searched indices 0..{current_index}. \
384 The key may have been generated outside this provider."
385 )))
386 }
387
388 fn get_cached_pks(&self) -> Result<Vec<bitcoin::XOnlyPublicKey>, Error> {
389 let cache = self
390 .key_cache
391 .read()
392 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
393
394 Ok(cache.keys().copied().collect())
395 }
396}
397
398impl DiscoverableKeyProvider for Bip32KeyProvider {
399 fn get_derivation_index_for_pk(&self, pk: &bitcoin::XOnlyPublicKey) -> Option<u32> {
400 let cache = self.key_cache.read().ok()?;
401 cache.get(pk).map(|v| v.path_index)
402 }
403
404 fn derive_at_discovery_index(&self, index: u32) -> Result<Option<Keypair>, Error> {
405 self.derive_at_index(index).map(Some)
406 }
407
408 fn cache_discovered_keypair(&self, index: u32, kp: Keypair) -> Result<(), Error> {
409 let pk = kp.x_only_public_key().0;
410
411 {
413 let mut cache = self
414 .key_cache
415 .write()
416 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
417 cache.insert(
418 pk,
419 KeyCacheValue {
420 path_index: index,
421 kp,
422 used: true,
423 },
424 );
425 }
426
427 {
429 let mut next = self
430 .next_index
431 .lock()
432 .map_err(|e| Error::ad_hoc(format!("Failed to lock next_index: {e}")))?;
433 if index >= *next {
434 *next = index
435 .checked_add(1)
436 .ok_or_else(|| Error::ad_hoc("Key derivation index overflow"))?;
437 }
438 }
439
440 Ok(())
441 }
442
443 fn cache_keypair_at_index(&self, index: u32) -> Result<(), Error> {
444 let kp = self.derive_at_index(index)?;
445 let pk = kp.x_only_public_key().0;
446 let mut cache = self
447 .key_cache
448 .write()
449 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
450 cache.insert(
451 pk,
452 KeyCacheValue {
453 path_index: index,
454 kp,
455 used: true,
456 },
457 );
458 Ok(())
459 }
460
461 fn mark_as_used(&self, pk: &bitcoin::XOnlyPublicKey) -> Result<(), Error> {
462 {
464 let maybe_kp = {
465 let cache = self
466 .key_cache
467 .read()
468 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
469 cache.get(pk).copied()
470 };
471
472 match maybe_kp {
473 Some(KeyCacheValue {
474 path_index,
475 kp,
476 used: false,
477 }) => {
478 let mut cache = self
479 .key_cache
480 .write()
481 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
482 cache.insert(
483 *pk,
484 KeyCacheValue {
485 path_index,
486 kp,
487 used: true,
488 },
489 );
490 return Ok(());
491 }
492 Some(KeyCacheValue { used: true, .. }) => {
493 return Ok(());
495 }
496 _ => {
497 }
499 }
500 }
501
502 let current_index = {
504 let next_index = self
505 .next_index
506 .lock()
507 .map_err(|e| Error::ad_hoc(format!("Failed to lock next_index: {e}")))?;
508 *next_index
509 };
510
511 for i in 0..current_index {
513 let kp = self.derive_at_index(i)?;
514 let derived_pk = kp.x_only_public_key().0;
515
516 if &derived_pk == pk {
517 let mut cache = self
519 .key_cache
520 .write()
521 .map_err(|e| Error::ad_hoc(format!("Failed to lock key_cache: {e}")))?;
522 cache.insert(
523 derived_pk,
524 KeyCacheValue {
525 path_index: i,
526 kp,
527 used: true,
528 },
529 );
530 return Ok(());
531 }
532 }
533
534 Err(Error::ad_hoc(format!(
535 "Public key {pk} not found in HD wallet. \
536 Searched indices 0..{current_index}. \
537 The key may have been generated outside this provider."
538 )))
539 }
540}
541
542impl<T: KeyProvider + ?Sized> KeyProvider for Arc<T> {
544 fn get_next_keypair(&self, keypair_index: KeypairIndex) -> Result<Keypair, Error> {
545 (**self).get_next_keypair(keypair_index)
546 }
547
548 fn get_keypair_for_path(&self, path: &[u32]) -> Result<Keypair, Error> {
549 (**self).get_keypair_for_path(path)
550 }
551
552 fn get_keypair_for_pk(&self, pk: &bitcoin::XOnlyPublicKey) -> Result<Keypair, Error> {
553 (**self).get_keypair_for_pk(pk)
554 }
555
556 fn get_cached_pks(&self) -> Result<Vec<bitcoin::XOnlyPublicKey>, Error> {
557 (**self).get_cached_pks()
558 }
559}
560
561impl<T: DiscoverableKeyProvider + ?Sized> DiscoverableKeyProvider for Arc<T> {
562 fn get_derivation_index_for_pk(&self, pk: &bitcoin::XOnlyPublicKey) -> Option<u32> {
563 (**self).get_derivation_index_for_pk(pk)
564 }
565
566 fn derive_at_discovery_index(&self, index: u32) -> Result<Option<Keypair>, Error> {
567 (**self).derive_at_discovery_index(index)
568 }
569
570 fn cache_discovered_keypair(&self, index: u32, kp: Keypair) -> Result<(), Error> {
571 (**self).cache_discovered_keypair(index, kp)
572 }
573
574 fn cache_keypair_at_index(&self, index: u32) -> Result<(), Error> {
575 (**self).cache_keypair_at_index(index)
576 }
577
578 fn mark_as_used(&self, pk: &bitcoin::XOnlyPublicKey) -> Result<(), Error> {
579 (**self).mark_as_used(pk)
580 }
581}
582
583#[cfg(test)]
584mod tests {
585 use super::*;
586 use bitcoin::Network;
587 use std::str::FromStr;
588
589 #[test]
590 fn cache_keypair_at_index_hydrates_hd_lookup_after_restart() {
591 let seed = [7_u8; 32];
592 let master = Xpriv::new_master(Network::Regtest, &seed).unwrap();
593 let base_path = DerivationPath::from_str("m/86'/1'/0'/0").unwrap();
594
595 let original = Bip32KeyProvider::new(master, base_path.clone());
596 let expected = original.derive_at_discovery_index(7).unwrap().unwrap();
597 let expected_pk = expected.x_only_public_key().0;
598 let first_receive_pk = original
599 .derive_at_discovery_index(0)
600 .unwrap()
601 .unwrap()
602 .x_only_public_key()
603 .0;
604
605 let restarted = Bip32KeyProvider::new(master, base_path);
606 assert!(restarted.get_keypair_for_pk(&expected_pk).is_err());
607
608 restarted.cache_keypair_at_index(7).unwrap();
609 let actual = restarted.get_keypair_for_pk(&expected_pk).unwrap();
610
611 assert_eq!(actual.x_only_public_key().0, expected_pk);
612 assert_eq!(restarted.get_derivation_index_for_pk(&expected_pk), Some(7));
613 assert_eq!(
614 restarted
615 .get_next_keypair(KeypairIndex::LastUnused)
616 .unwrap()
617 .x_only_public_key()
618 .0,
619 first_receive_pk
620 );
621 }
622}