use alloc::string::String;
use alloc::vec::Vec;
use kobe_primitives::slip10::DerivedKey;
use kobe_primitives::{Derive, DerivedAccount, Wallet};
use zeroize::Zeroizing;
use crate::DeriveError;
use crate::derivation_style::DerivationStyle;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct DerivedAddress {
pub path: String,
pub private_key_hex: Zeroizing<String>,
pub keypair_base58: Zeroizing<String>,
pub public_key_hex: String,
pub address: String,
}
impl DerivedAddress {
pub fn private_key_bytes(&self) -> Result<Zeroizing<[u8; 32]>, DeriveError> {
let mut buf = Zeroizing::new([0u8; 32]);
hex::decode_to_slice(self.private_key_hex.as_str(), buf.as_mut_slice()).map_err(|e| {
kobe_primitives::DeriveError::InvalidHex(alloc::format!("private_key_hex: {e}"))
})?;
Ok(buf)
}
pub fn public_key_bytes(&self) -> Result<Vec<u8>, DeriveError> {
hex::decode(&self.public_key_hex)
.map_err(|e| {
kobe_primitives::DeriveError::InvalidHex(alloc::format!("public_key_hex: {e}"))
})
.map_err(Into::into)
}
}
#[derive(Debug)]
pub struct Deriver<'a> {
wallet: &'a Wallet,
}
impl<'a> Deriver<'a> {
#[inline]
#[must_use]
pub const fn new(wallet: &'a Wallet) -> Self {
Self { wallet }
}
#[inline]
pub fn derive(&self, index: u32) -> Result<DerivedAddress, DeriveError> {
self.derive_with(DerivationStyle::Standard, index)
}
pub fn derive_with(
&self,
style: DerivationStyle,
index: u32,
) -> Result<DerivedAddress, DeriveError> {
let path = style.path(index);
let derived = DerivedKey::derive_path(self.wallet.seed(), &path)?;
Ok(build_derived_address(&derived, path))
}
#[inline]
pub fn derive_many(&self, start: u32, count: u32) -> Result<Vec<DerivedAddress>, DeriveError> {
self.derive_many_with(DerivationStyle::Standard, start, count)
}
pub fn derive_many_with(
&self,
style: DerivationStyle,
start: u32,
count: u32,
) -> Result<Vec<DerivedAddress>, DeriveError> {
let end = start
.checked_add(count)
.ok_or(kobe_primitives::DeriveError::IndexOverflow)?;
(start..end)
.map(|index| self.derive_with(style, index))
.collect()
}
pub fn derive_path(&self, path: &str) -> Result<DerivedAddress, DeriveError> {
let derived = DerivedKey::derive_path(self.wallet.seed(), path)?;
Ok(build_derived_address(&derived, String::from(path)))
}
}
impl Derive for Deriver<'_> {
type Error = DeriveError;
fn derive(&self, index: u32) -> Result<DerivedAccount, DeriveError> {
let da = self.derive_with(DerivationStyle::Standard, index)?;
Ok(DerivedAccount::new(
da.path,
da.private_key_hex,
da.public_key_hex,
da.address,
))
}
fn derive_path(&self, path: &str) -> Result<DerivedAccount, DeriveError> {
let da = Deriver::derive_path(self, path)?;
Ok(DerivedAccount::new(
da.path,
da.private_key_hex,
da.public_key_hex,
da.address,
))
}
}
fn build_derived_address(derived: &DerivedKey, path: String) -> DerivedAddress {
let signing_key = derived.to_signing_key();
let verifying_key = signing_key.verifying_key();
let public_key_bytes = verifying_key.as_bytes();
let mut keypair_bytes = Zeroizing::new([0u8; 64]);
let (left, right) = keypair_bytes.split_at_mut(32);
left.copy_from_slice(derived.private_key.as_slice());
right.copy_from_slice(public_key_bytes);
let keypair_b58 = bs58::encode(&*keypair_bytes).into_string();
DerivedAddress {
path,
private_key_hex: Zeroizing::new(hex::encode(derived.private_key.as_slice())),
keypair_base58: Zeroizing::new(keypair_b58),
public_key_hex: hex::encode(public_key_bytes),
address: bs58::encode(public_key_bytes).into_string(),
}
}
#[cfg(test)]
#[allow(clippy::indexing_slicing, reason = "test assertions")]
mod tests {
use super::*;
fn test_wallet() -> Wallet {
Wallet::from_mnemonic(
"abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about",
None,
)
.unwrap()
}
#[test]
fn test_derive_address() {
let wallet = test_wallet();
let deriver = Deriver::new(&wallet);
let addr = deriver.derive(0).unwrap();
assert!(addr.address.len() >= 32 && addr.address.len() <= 44);
assert_eq!(addr.path, "m/44'/501'/0'/0'");
}
#[test]
fn test_derive_many() {
let wallet = test_wallet();
let deriver = Deriver::new(&wallet);
let addresses = deriver.derive_many(0, 3).unwrap();
assert_eq!(addresses.len(), 3);
assert_eq!(addresses[0].path, "m/44'/501'/0'/0'");
assert_eq!(addresses[1].path, "m/44'/501'/1'/0'");
assert_eq!(addresses[2].path, "m/44'/501'/2'/0'");
assert_ne!(addresses[0].address, addresses[1].address);
assert_ne!(addresses[1].address, addresses[2].address);
}
#[test]
fn test_deterministic_derivation() {
let wallet = test_wallet();
let deriver = Deriver::new(&wallet);
let addr1 = deriver.derive(0).unwrap();
let addr2 = deriver.derive(0).unwrap();
assert_eq!(addr1.address, addr2.address);
assert_eq!(*addr1.private_key_hex, *addr2.private_key_hex);
}
#[test]
fn test_derive_with_trust() {
let wallet = test_wallet();
let deriver = Deriver::new(&wallet);
let addr = deriver.derive_with(DerivationStyle::Trust, 0).unwrap();
assert_eq!(addr.path, "m/44'/501'/0'");
assert!(addr.address.len() >= 32 && addr.address.len() <= 44);
}
#[test]
fn test_derive_with_ledger_live() {
let wallet = test_wallet();
let deriver = Deriver::new(&wallet);
let addr = deriver.derive_with(DerivationStyle::LedgerLive, 0).unwrap();
assert_eq!(addr.path, "m/44'/501'/0'/0'/0'");
assert!(addr.address.len() >= 32 && addr.address.len() <= 44);
}
#[test]
fn test_different_styles_produce_different_addresses() {
let wallet = test_wallet();
let deriver = Deriver::new(&wallet);
let standard = deriver.derive_with(DerivationStyle::Standard, 0).unwrap();
let trust = deriver.derive_with(DerivationStyle::Trust, 0).unwrap();
let ledger_live = deriver.derive_with(DerivationStyle::LedgerLive, 0).unwrap();
let legacy = deriver.derive_with(DerivationStyle::Legacy, 0).unwrap();
assert_ne!(standard.address, trust.address);
assert_ne!(standard.address, ledger_live.address);
assert_ne!(standard.address, legacy.address);
assert_ne!(trust.address, ledger_live.address);
}
#[test]
fn kat_solana_standard_index0() {
let wallet = test_wallet();
let addr = Deriver::new(&wallet).derive(0).unwrap();
assert_eq!(addr.address, "HAgk14JpMQLgt6rVgv7cBQFJWFto5Dqxi472uT3DKpqk");
}
#[test]
fn bytes_accessors_roundtrip() {
let wallet = test_wallet();
let addr = Deriver::new(&wallet).derive(0).unwrap();
let sk = addr.private_key_bytes().unwrap();
assert_eq!(sk.len(), 32);
assert_eq!(hex::encode(*sk), addr.private_key_hex.as_str());
let pk = addr.public_key_bytes().unwrap();
assert_eq!(pk.len(), 32);
assert_eq!(hex::encode(&pk), addr.public_key_hex);
}
}