use bitflags::bitflags;
use transparent::address::TransparentAddress::*;
use zcash_address::{
ConversionError, TryFromAddress,
unified::{Container, Receiver},
};
use zcash_client_backend::data_api::AccountSource;
use zcash_keys::{
address::{Address, UnifiedAddress},
keys::{ReceiverRequirement, ReceiverRequirements},
};
use zcash_protocol::{PoolType, ShieldedPool, consensus::NetworkType, memo::MemoBytes};
use zip32::DiversifierIndex;
#[cfg(feature = "transparent-inputs")]
use {
super::transparent::SchedulingError,
std::time::{Duration, SystemTime},
transparent::keys::TransparentKeyScope,
zcash_client_backend::data_api::TransparentKeyOrigin,
zcash_keys::keys::AddressGenerationError,
};
use crate::error::SqliteClientError;
#[cfg(feature = "zcashd-compat")]
use zcash_keys::keys::zcashd;
pub(crate) const LEGACY_ADDRESS_INDEX_NULL: i64 = -1;
pub(crate) fn pool_code(pool_type: PoolType) -> i64 {
match pool_type {
PoolType::Transparent => 0i64,
PoolType::Shielded(ShieldedPool::Sapling) => 2i64,
PoolType::Shielded(ShieldedPool::Orchard) => 3i64,
PoolType::Shielded(ShieldedPool::Ironwood) => 4i64,
}
}
pub(crate) fn parse_pool_code(code: i64) -> Result<PoolType, SqliteClientError> {
match code {
0i64 => Ok(PoolType::Transparent),
2i64 => Ok(PoolType::SAPLING),
3i64 => Ok(PoolType::ORCHARD),
4i64 => Ok(PoolType::IRONWOOD),
_ => Err(SqliteClientError::CorruptedData(format!(
"Invalid pool code: {code}"
))),
}
}
pub(crate) fn account_kind_code(value: &AccountSource) -> u32 {
match value {
AccountSource::Derived { .. } => 0,
AccountSource::Imported { .. } => 1,
}
}
pub(crate) fn encode_diversifier_index_be(idx: DiversifierIndex) -> [u8; 11] {
let mut di_be = *idx.as_bytes();
di_be.reverse();
di_be
}
pub(crate) fn decode_diversifier_index_be(
di_be: Option<Vec<u8>>,
) -> Result<Option<DiversifierIndex>, SqliteClientError> {
di_be
.map(|di_be_bytes| {
let mut di_be: [u8; 11] = di_be_bytes.try_into().map_err(|_| {
SqliteClientError::CorruptedData(
"Diversifier index is not an 11-byte value".to_owned(),
)
})?;
di_be.reverse();
Ok(DiversifierIndex::from(di_be))
})
.transpose()
}
pub(crate) fn memo_repr(memo: Option<&MemoBytes>) -> Option<&[u8]> {
memo.map(|m| {
if m == &MemoBytes::empty() {
&[0xf6]
} else {
m.as_slice()
}
})
}
#[cfg(feature = "transparent-inputs")]
pub(crate) fn epoch_seconds(t: SystemTime) -> Result<i64, SchedulingError> {
let integer_seconds_since_epoch =
i64::try_from(t.duration_since(SystemTime::UNIX_EPOCH)?.as_secs())?;
Ok(integer_seconds_since_epoch)
}
#[cfg(feature = "transparent-inputs")]
pub(crate) fn decode_epoch_seconds(i: i64) -> Result<SystemTime, SchedulingError> {
Ok(SystemTime::UNIX_EPOCH + Duration::from_secs(u64::try_from(i)?))
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub(crate) enum KeyScope {
Zip32(zip32::Scope),
Ephemeral,
Foreign,
}
impl KeyScope {
pub(crate) const EXTERNAL: KeyScope = KeyScope::Zip32(zip32::Scope::External);
pub(crate) const INTERNAL: KeyScope = KeyScope::Zip32(zip32::Scope::Internal);
pub(crate) fn encode(&self) -> i64 {
match self {
KeyScope::Zip32(zip32::Scope::External) => 0i64,
KeyScope::Zip32(zip32::Scope::Internal) => 1i64,
KeyScope::Ephemeral => 2i64,
KeyScope::Foreign => -1i64,
}
}
pub(crate) fn decode(code: i64) -> Result<Self, SqliteClientError> {
match code {
0i64 => Ok(KeyScope::EXTERNAL),
1i64 => Ok(KeyScope::INTERNAL),
2i64 => Ok(KeyScope::Ephemeral),
-1i64 => Ok(KeyScope::Foreign),
other => Err(SqliteClientError::CorruptedData(format!(
"Invalid key scope code: {other}"
))),
}
}
#[cfg(feature = "transparent-inputs")]
pub(crate) fn as_transparent(&self) -> Option<TransparentKeyScope> {
match self {
KeyScope::Zip32(scope) => Some(TransparentKeyScope::from(*scope)),
KeyScope::Ephemeral => Some(TransparentKeyScope::custom(2).expect("valid scope")),
KeyScope::Foreign => None,
}
}
#[cfg(feature = "transparent-inputs")]
pub(crate) fn as_key_origin(&self) -> TransparentKeyOrigin {
match self {
KeyScope::Zip32(scope) => TransparentKeyOrigin::Derived {
scope: TransparentKeyScope::from(*scope),
},
KeyScope::Ephemeral => TransparentKeyOrigin::Derived {
scope: TransparentKeyScope::custom(2).expect("valid scope"),
},
KeyScope::Foreign => TransparentKeyOrigin::Imported,
}
}
}
impl From<zip32::Scope> for KeyScope {
fn from(value: zip32::Scope) -> Self {
KeyScope::Zip32(value)
}
}
#[cfg(feature = "transparent-inputs")]
impl From<KeyScope> for Option<TransparentKeyScope> {
fn from(value: KeyScope) -> Self {
value.as_transparent()
}
}
#[cfg(feature = "transparent-inputs")]
impl TryFrom<TransparentKeyScope> for KeyScope {
type Error = AddressGenerationError;
fn try_from(value: TransparentKeyScope) -> Result<Self, Self::Error> {
match value {
TransparentKeyScope::EXTERNAL => Ok(KeyScope::EXTERNAL),
TransparentKeyScope::INTERNAL => Ok(KeyScope::INTERNAL),
TransparentKeyScope::EPHEMERAL => Ok(KeyScope::Ephemeral),
_ => Err(AddressGenerationError::UnsupportedTransparentKeyScope(
value,
)),
}
}
}
impl TryFrom<KeyScope> for zip32::Scope {
type Error = ();
fn try_from(value: KeyScope) -> Result<Self, Self::Error> {
match value {
KeyScope::Zip32(scope) => Ok(scope),
KeyScope::Ephemeral | KeyScope::Foreign => Err(()),
}
}
}
bitflags! {
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct ReceiverFlags: i64 {
const UNKNOWN = 0b00000000;
const P2PKH = 0b00000001;
const P2SH = 0b00000010;
const SAPLING = 0b00000100;
const ORCHARD = 0b00001000;
}
}
impl ReceiverFlags {
pub(crate) fn required(request: ReceiverRequirements) -> Self {
let mut flags = ReceiverFlags::UNKNOWN;
if matches!(request.orchard(), ReceiverRequirement::Require) {
flags |= ReceiverFlags::ORCHARD;
}
if matches!(request.sapling(), ReceiverRequirement::Require) {
flags |= ReceiverFlags::SAPLING;
}
if matches!(request.p2pkh(), ReceiverRequirement::Require) {
flags |= ReceiverFlags::P2PKH;
}
flags
}
pub(crate) fn omitted(request: ReceiverRequirements) -> Self {
let mut flags = ReceiverFlags::UNKNOWN;
if matches!(request.orchard(), ReceiverRequirement::Omit) {
flags |= ReceiverFlags::ORCHARD;
}
if matches!(request.sapling(), ReceiverRequirement::Omit) {
flags |= ReceiverFlags::SAPLING;
}
if matches!(request.p2pkh(), ReceiverRequirement::Omit) {
flags |= ReceiverFlags::P2PKH;
}
flags
}
}
impl From<&UnifiedAddress> for ReceiverFlags {
fn from(value: &UnifiedAddress) -> Self {
let mut flags = ReceiverFlags::UNKNOWN;
match value.transparent() {
Some(PublicKeyHash(_)) => {
flags |= ReceiverFlags::P2PKH;
}
Some(ScriptHash(_)) => {
flags |= ReceiverFlags::P2SH;
}
_ => {}
}
if value.has_sapling() {
flags |= ReceiverFlags::SAPLING;
}
if value.has_orchard() {
flags |= ReceiverFlags::ORCHARD;
}
flags
}
}
impl From<&Address> for ReceiverFlags {
fn from(address: &Address) -> Self {
match address {
Address::Sapling(_) => ReceiverFlags::SAPLING,
Address::Transparent(addr) => match addr {
PublicKeyHash(_) => ReceiverFlags::P2PKH,
ScriptHash(_) => ReceiverFlags::P2SH,
},
Address::Unified(ua) => ReceiverFlags::from(ua),
Address::Tex(_) => ReceiverFlags::P2PKH,
}
}
}
impl TryFromAddress for ReceiverFlags {
type Error = ();
fn try_from_sapling(
_net: NetworkType,
_data: [u8; 43],
) -> Result<Self, ConversionError<Self::Error>> {
Ok(ReceiverFlags::SAPLING)
}
fn try_from_unified(
_net: NetworkType,
data: zcash_address::unified::Address,
) -> Result<Self, ConversionError<Self::Error>> {
let mut result = ReceiverFlags::UNKNOWN;
for i in data.items() {
match i {
Receiver::Orchard(_) => result |= ReceiverFlags::ORCHARD,
Receiver::Sapling(_) => result |= ReceiverFlags::SAPLING,
Receiver::P2pkh(_) => result |= ReceiverFlags::P2PKH,
Receiver::P2sh(_) => result |= ReceiverFlags::P2SH,
Receiver::Unknown { .. } => {}
}
}
Ok(result)
}
fn try_from_transparent_p2pkh(
_net: NetworkType,
_data: [u8; 20],
) -> Result<Self, ConversionError<Self::Error>> {
Ok(ReceiverFlags::P2PKH)
}
fn try_from_transparent_p2sh(
_net: NetworkType,
_data: [u8; 20],
) -> Result<Self, ConversionError<Self::Error>> {
Ok(ReceiverFlags::P2SH)
}
fn try_from_tex(
_net: NetworkType,
_data: [u8; 20],
) -> Result<Self, ConversionError<Self::Error>> {
Ok(ReceiverFlags::P2PKH)
}
}
#[cfg(feature = "zcashd-compat")]
pub(crate) fn decode_legacy_account_index(
legacy_account_index: i64,
) -> Result<Option<zcashd::LegacyAddressIndex>, SqliteClientError> {
match legacy_account_index {
LEGACY_ADDRESS_INDEX_NULL => Ok(None),
_ => u32::try_from(legacy_account_index)
.map_err(|_| ())
.and_then(zcashd::LegacyAddressIndex::try_from)
.map(Some)
.map_err(|_| {
SqliteClientError::CorruptedData(
"Legacy zcashd address index is out of range.".to_string(),
)
}),
}
}
#[cfg(feature = "zcashd-compat")]
pub(crate) fn encode_legacy_account_index(
legacy_account_index: Option<zcashd::LegacyAddressIndex>,
) -> i64 {
legacy_account_index
.map(u32::from)
.map_or(LEGACY_ADDRESS_INDEX_NULL, i64::from)
}
#[cfg(test)]
mod tests {
use zcash_protocol::{PoolType, ShieldedPool};
use super::{parse_pool_code, pool_code};
#[test]
fn pool_code_round_trips() {
for pool in [
PoolType::Transparent,
PoolType::Shielded(ShieldedPool::Sapling),
PoolType::Shielded(ShieldedPool::Orchard),
PoolType::Shielded(ShieldedPool::Ironwood),
] {
assert_eq!(parse_pool_code(pool_code(pool)).unwrap(), pool);
}
}
#[test]
fn ironwood_pool_code_is_distinct() {
let codes = [
pool_code(PoolType::Transparent),
pool_code(PoolType::Shielded(ShieldedPool::Sapling)),
pool_code(PoolType::Shielded(ShieldedPool::Orchard)),
pool_code(PoolType::Shielded(ShieldedPool::Ironwood)),
];
let unique: std::collections::BTreeSet<_> = codes.iter().collect();
assert_eq!(unique.len(), codes.len(), "pool codes must be distinct");
}
}