use alloc::collections::BTreeMap;
use alloc::string::String;
use alloc::vec::Vec;
use core::cmp::Ordering;
use getset::Getters;
use serde::{Deserialize, Serialize};
use serde_with::serde_as;
use crate::{
common::{Global, Zip32Derivation},
roles::combiner::{merge_map, merge_optional},
};
const GROTH_PROOF_SIZE: usize = 48 + 96 + 48;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Getters)]
pub struct Bundle {
#[getset(get = "pub")]
pub(crate) spends: Vec<Spend>,
#[getset(get = "pub")]
pub(crate) outputs: Vec<Output>,
#[getset(get = "pub")]
pub(crate) value_sum: i128,
#[getset(get = "pub")]
pub(crate) anchor: Option<[u8; 32]>,
pub(crate) bsk: Option<[u8; 32]>,
}
pub(crate) const EMPTY_BUNDLE: Bundle = Bundle {
spends: Vec::new(),
outputs: Vec::new(),
value_sum: 0,
anchor: None,
bsk: None,
};
pub(crate) const DEFAULT_ANCHOR: [u8; 32] = [0; 32];
#[serde_as]
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Getters)]
pub struct Spend {
#[getset(get = "pub")]
pub(crate) cv: [u8; 32],
#[getset(get = "pub")]
pub(crate) nullifier: [u8; 32],
#[getset(get = "pub")]
pub(crate) rk: [u8; 32],
#[serde_as(as = "Option<[_; GROTH_PROOF_SIZE]>")]
pub(crate) zkproof: Option<[u8; GROTH_PROOF_SIZE]>,
#[serde_as(as = "Option<[_; 64]>")]
pub(crate) spend_auth_sig: Option<[u8; 64]>,
#[serde_as(as = "Option<[_; 43]>")]
pub(crate) recipient: Option<[u8; 43]>,
pub(crate) value: Option<u64>,
pub(crate) rcm: Option<[u8; 32]>,
pub(crate) rseed: Option<[u8; 32]>,
pub(crate) rcv: Option<[u8; 32]>,
pub(crate) proof_generation_key: Option<([u8; 32], [u8; 32])>,
pub(crate) witness: Option<(u32, [[u8; 32]; 32])>,
pub(crate) alpha: Option<[u8; 32]>,
pub(crate) zip32_derivation: Option<Zip32Derivation>,
pub(crate) dummy_ask: Option<[u8; 32]>,
#[getset(get = "pub")]
pub(crate) proprietary: BTreeMap<String, Vec<u8>>,
}
#[serde_as]
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Getters)]
pub struct Output {
#[getset(get = "pub")]
pub(crate) cv: [u8; 32],
#[getset(get = "pub")]
pub(crate) cmu: [u8; 32],
#[getset(get = "pub")]
pub(crate) ephemeral_key: [u8; 32],
#[getset(get = "pub")]
pub(crate) enc_ciphertext: Vec<u8>,
#[getset(get = "pub")]
pub(crate) out_ciphertext: Vec<u8>,
#[serde_as(as = "Option<[_; GROTH_PROOF_SIZE]>")]
pub(crate) zkproof: Option<[u8; GROTH_PROOF_SIZE]>,
#[serde_as(as = "Option<[_; 43]>")]
#[getset(get = "pub")]
pub(crate) recipient: Option<[u8; 43]>,
#[getset(get = "pub")]
pub(crate) value: Option<u64>,
#[getset(get = "pub")]
pub(crate) rseed: Option<[u8; 32]>,
pub(crate) rcv: Option<[u8; 32]>,
pub(crate) ock: Option<[u8; 32]>,
pub(crate) zip32_derivation: Option<Zip32Derivation>,
#[getset(get = "pub")]
pub(crate) user_address: Option<String>,
#[getset(get = "pub")]
pub(crate) proprietary: BTreeMap<String, Vec<u8>>,
}
impl Bundle {
pub(crate) fn merge(
mut self,
other: Self,
self_global: &Global,
other_global: &Global,
) -> Option<Self> {
let Self {
mut spends,
mut outputs,
value_sum,
anchor,
bsk,
} = other;
match (self.bsk.as_mut(), bsk) {
(Some(lhs), Some(rhs)) if lhs != &rhs => return None,
(Some(_), _) | (_, Some(_))
if self.spends.len() != spends.len()
|| self.outputs.len() != outputs.len()
|| self.value_sum != value_sum =>
{
return None;
}
(Some(_), _) | (_, Some(_)) => (),
(None, None) => {
let (spends_cmp_other, outputs_cmp_other) = match (
self.spends.len().cmp(&spends.len()),
self.outputs.len().cmp(&outputs.len()),
) {
(Ordering::Less, Ordering::Greater) | (Ordering::Greater, Ordering::Less) => {
return None;
}
(spends, outputs) => (spends, outputs),
};
match (
self_global.shielded_modifiable(),
other_global.shielded_modifiable(),
spends_cmp_other,
) {
(false, _, Ordering::Less) | (_, false, Ordering::Greater) => return None,
(true, _, Ordering::Less) => {
self.spends.extend(spends.drain(self.spends.len()..))
}
(_, _, Ordering::Equal) | (_, true, Ordering::Greater) => (),
}
match (
self_global.shielded_modifiable(),
other_global.shielded_modifiable(),
outputs_cmp_other,
) {
(false, _, Ordering::Less) | (_, false, Ordering::Greater) => return None,
(true, _, Ordering::Less) => {
self.outputs.extend(outputs.drain(self.outputs.len()..))
}
(_, _, Ordering::Equal) | (_, true, Ordering::Greater) => (),
}
if matches!(spends_cmp_other, Ordering::Less)
|| matches!(outputs_cmp_other, Ordering::Less)
{
self.value_sum = value_sum;
}
}
}
if !merge_optional(&mut self.anchor, anchor) {
return None;
}
for (lhs, rhs) in self.spends.iter_mut().zip(spends) {
let Spend {
cv,
nullifier,
rk,
zkproof,
spend_auth_sig,
recipient,
value,
rcm,
rseed,
rcv,
proof_generation_key,
witness,
alpha,
zip32_derivation,
dummy_ask,
proprietary,
} = rhs;
if lhs.cv != cv || lhs.nullifier != nullifier || lhs.rk != rk {
return None;
}
if !(merge_optional(&mut lhs.zkproof, zkproof)
&& merge_optional(&mut lhs.spend_auth_sig, spend_auth_sig)
&& merge_optional(&mut lhs.recipient, recipient)
&& merge_optional(&mut lhs.value, value)
&& merge_optional(&mut lhs.rcm, rcm)
&& merge_optional(&mut lhs.rseed, rseed)
&& merge_optional(&mut lhs.rcv, rcv)
&& merge_optional(&mut lhs.proof_generation_key, proof_generation_key)
&& merge_optional(&mut lhs.witness, witness)
&& merge_optional(&mut lhs.alpha, alpha)
&& merge_optional(&mut lhs.zip32_derivation, zip32_derivation)
&& merge_optional(&mut lhs.dummy_ask, dummy_ask)
&& merge_map(&mut lhs.proprietary, proprietary))
{
return None;
}
}
for (lhs, rhs) in self.outputs.iter_mut().zip(outputs) {
let Output {
cv,
cmu,
ephemeral_key,
enc_ciphertext,
out_ciphertext,
zkproof,
recipient,
value,
rseed,
rcv,
ock,
zip32_derivation,
user_address,
proprietary,
} = rhs;
if lhs.cv != cv
|| lhs.cmu != cmu
|| lhs.ephemeral_key != ephemeral_key
|| lhs.enc_ciphertext != enc_ciphertext
|| lhs.out_ciphertext != out_ciphertext
{
return None;
}
if !(merge_optional(&mut lhs.zkproof, zkproof)
&& merge_optional(&mut lhs.recipient, recipient)
&& merge_optional(&mut lhs.value, value)
&& merge_optional(&mut lhs.rseed, rseed)
&& merge_optional(&mut lhs.rcv, rcv)
&& merge_optional(&mut lhs.ock, ock)
&& merge_optional(&mut lhs.zip32_derivation, zip32_derivation)
&& merge_optional(&mut lhs.user_address, user_address)
&& merge_map(&mut lhs.proprietary, proprietary))
{
return None;
}
}
Some(self)
}
}
pub(crate) mod v1 {
use alloc::vec::Vec;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct Bundle {
spends: Vec<super::Spend>,
outputs: Vec<super::Output>,
value_sum: i128,
anchor: [u8; 32],
bsk: Option<[u8; 32]>,
}
impl TryFrom<super::Bundle> for Bundle {
type Error = crate::EncodingError;
fn try_from(bundle: super::Bundle) -> Result<Self, Self::Error> {
let anchor = match bundle.anchor {
Some(anchor) => anchor,
None if bundle.spends.is_empty() => super::DEFAULT_ANCHOR,
None => return Err(crate::EncodingError::RequiresV2),
};
Ok(Self {
spends: bundle.spends,
outputs: bundle.outputs,
value_sum: bundle.value_sum,
anchor,
bsk: bundle.bsk,
})
}
}
impl From<Bundle> for super::Bundle {
fn from(bundle: Bundle) -> Self {
Self {
spends: bundle.spends,
outputs: bundle.outputs,
value_sum: bundle.value_sum,
anchor: Some(bundle.anchor),
bsk: bundle.bsk,
}
}
}
}
pub(crate) mod v2 {
pub(crate) fn encode(bundle: super::Bundle) -> Option<super::Bundle> {
(!is_default_empty(&bundle)).then_some(bundle)
}
fn is_default_empty(bundle: &super::Bundle) -> bool {
let mut bundle = bundle.clone();
if bundle.anchor == Some(super::DEFAULT_ANCHOR) {
bundle.anchor = None;
}
bundle == super::EMPTY_BUNDLE
}
}
#[cfg(feature = "sapling")]
impl Bundle {
pub(crate) fn into_parsed(
self,
anchor_requirement: crate::common::AnchorRequirement,
) -> Result<Parsed, ParseError> {
let wire_anchor = self.anchor;
let anchor = anchor_requirement
.resolve(
wire_anchor,
self.spends.is_empty() && self.outputs.is_empty(),
)
.ok_or(ParseError::MissingAnchor)?;
let spends = self
.spends
.into_iter()
.map(|spend| {
sapling::pczt::Spend::parse(
spend.cv,
spend.nullifier,
spend.rk,
spend.zkproof,
spend.spend_auth_sig,
spend.recipient,
spend.value,
spend.rcm,
spend.rseed,
spend.rcv,
spend.proof_generation_key,
spend.witness,
spend.alpha,
spend
.zip32_derivation
.map(|z| {
sapling::pczt::Zip32Derivation::parse(
z.seed_fingerprint,
z.derivation_path,
)
})
.transpose()?,
spend.dummy_ask,
spend.proprietary,
)
})
.collect::<Result<_, _>>()?;
let outputs = self
.outputs
.into_iter()
.map(|output| {
sapling::pczt::Output::parse(
output.cv,
output.cmu,
output.ephemeral_key,
output.enc_ciphertext,
output.out_ciphertext,
output.zkproof,
output.recipient,
output.value,
output.rseed,
output.rcv,
output.ock,
output
.zip32_derivation
.map(|z| {
sapling::pczt::Zip32Derivation::parse(
z.seed_fingerprint,
z.derivation_path,
)
})
.transpose()?,
output.user_address,
output.proprietary,
)
})
.collect::<Result<_, _>>()?;
let bundle =
sapling::pczt::Bundle::parse(spends, outputs, self.value_sum, anchor, self.bsk)?;
Ok(Parsed {
bundle,
wire_anchor,
})
}
pub(crate) fn serialize_from(bundle: sapling::pczt::Bundle) -> Self {
let spends = bundle
.spends()
.iter()
.map(|spend| {
let (rcm, rseed) = match spend.rseed() {
Some(sapling::Rseed::BeforeZip212(rcm)) => (Some(rcm.to_bytes()), None),
Some(sapling::Rseed::AfterZip212(rseed)) => (None, Some(*rseed)),
None => (None, None),
};
Spend {
cv: spend.cv().to_bytes(),
nullifier: spend.nullifier().0,
rk: (*spend.rk()).into(),
zkproof: *spend.zkproof(),
spend_auth_sig: spend.spend_auth_sig().map(|s| s.into()),
recipient: spend.recipient().map(|recipient| recipient.to_bytes()),
value: spend.value().map(|value| value.inner()),
rcm,
rseed,
rcv: spend.rcv().as_ref().map(|rcv| rcv.inner().to_bytes()),
proof_generation_key: spend
.proof_generation_key()
.as_ref()
.map(|key| (key.ak.to_bytes(), key.nsk.to_bytes())),
witness: spend.witness().as_ref().map(|witness| {
(
u32::try_from(u64::from(witness.position()))
.expect("Sapling positions fit in u32"),
witness
.path_elems()
.iter()
.map(|node| node.to_bytes())
.collect::<Vec<_>>()[..]
.try_into()
.expect("path is length 32"),
)
}),
alpha: spend.alpha().map(|alpha| alpha.to_bytes()),
zip32_derivation: spend.zip32_derivation().as_ref().map(|z| Zip32Derivation {
seed_fingerprint: *z.seed_fingerprint(),
derivation_path: z.derivation_path().iter().map(|i| i.index()).collect(),
}),
dummy_ask: spend
.dummy_ask()
.as_ref()
.map(|dummy_ask| dummy_ask.to_bytes()),
proprietary: spend.proprietary().clone(),
}
})
.collect();
let outputs = bundle
.outputs()
.iter()
.map(|output| Output {
cv: output.cv().to_bytes(),
cmu: output.cmu().to_bytes(),
ephemeral_key: output.ephemeral_key().0,
enc_ciphertext: output.enc_ciphertext().to_vec(),
out_ciphertext: output.out_ciphertext().to_vec(),
zkproof: *output.zkproof(),
recipient: output.recipient().map(|recipient| recipient.to_bytes()),
value: output.value().map(|value| value.inner()),
rseed: *output.rseed(),
rcv: output.rcv().as_ref().map(|rcv| rcv.inner().to_bytes()),
ock: output.ock().as_ref().map(|ock| ock.0),
zip32_derivation: output.zip32_derivation().as_ref().map(|z| Zip32Derivation {
seed_fingerprint: *z.seed_fingerprint(),
derivation_path: z.derivation_path().iter().map(|i| i.index()).collect(),
}),
user_address: output.user_address().clone(),
proprietary: output.proprietary().clone(),
})
.collect();
Self {
spends,
outputs,
value_sum: bundle.value_sum().to_raw(),
anchor: Some(bundle.anchor().to_bytes()),
bsk: bundle.bsk().map(|bsk| bsk.into()),
}
}
}
#[cfg(feature = "sapling")]
#[derive(Debug)]
#[non_exhaustive]
pub enum ParseError {
MissingAnchor,
Bundle(sapling::pczt::ParseError),
}
#[cfg(feature = "sapling")]
impl From<sapling::pczt::ParseError> for ParseError {
fn from(e: sapling::pczt::ParseError) -> Self {
ParseError::Bundle(e)
}
}
#[cfg(all(feature = "sapling", feature = "prover"))]
#[derive(Debug)]
#[non_exhaustive]
pub enum AnchorConsistencyError {
IncompleteSpendData,
WitnessDoesNotRootToAnchor,
}
#[cfg(all(feature = "sapling", feature = "prover"))]
pub(crate) fn verify_witnesses_root_to_anchor(
bundle: &sapling::pczt::Bundle,
anchor: sapling::Anchor,
) -> Result<(), AnchorConsistencyError> {
for spend in bundle.spends() {
let Some(witness) = spend.witness() else {
continue;
};
let Some(value) = spend.value() else {
continue;
};
if value.inner() == 0 {
continue;
}
let recipient = spend
.recipient()
.ok_or(AnchorConsistencyError::IncompleteSpendData)?;
let rseed = (*spend.rseed()).ok_or(AnchorConsistencyError::IncompleteSpendData)?;
let note = sapling::Note::from_parts(recipient, *value, rseed);
let leaf = sapling::Node::from_cmu(¬e.cmu());
let computed_anchor: sapling::Anchor = witness.root(leaf).into();
if computed_anchor != anchor {
return Err(AnchorConsistencyError::WitnessDoesNotRootToAnchor);
}
}
Ok(())
}
#[cfg(feature = "sapling")]
pub(crate) struct Parsed {
pub(crate) bundle: sapling::pczt::Bundle,
pub(crate) wire_anchor: Option<[u8; 32]>,
}
#[cfg(feature = "sapling")]
impl Parsed {
pub(crate) fn reserialize(self) -> Bundle {
Bundle {
anchor: self.wire_anchor,
..Bundle::serialize_from(self.bundle)
}
}
}
#[cfg(all(test, feature = "sapling"))]
mod tests {
use rand_chacha::{ChaCha20Rng, rand_core::SeedableRng};
use super::{Bundle, ParseError};
use crate::common::AnchorRequirement;
fn output_only_bundle() -> Bundle {
let recipient = sapling::zip32::ExtendedSpendingKey::master(&[0; 32])
.to_diversifiable_full_viewing_key()
.default_address()
.1;
let mut builder = sapling::builder::Builder::new(
sapling::note_encryption::Zip212Enforcement::On,
sapling::builder::BundleType::DEFAULT,
sapling::Anchor::empty_tree(),
);
builder
.add_output(
None,
recipient,
sapling::value::NoteValue::from_raw(1),
[0; 512],
)
.unwrap();
let (bundle, _) = builder
.build_for_pczt(ChaCha20Rng::from_seed([0; 32]))
.unwrap();
let mut bundle = Bundle::serialize_from(bundle);
bundle.anchor = None;
bundle
}
#[test]
fn output_only_bundle_requires_anchor_when_required() {
assert!(matches!(
output_only_bundle().into_parsed(AnchorRequirement::Required),
Err(ParseError::MissingAnchor)
));
}
#[test]
fn output_only_bundle_preserves_absent_anchor_when_not_required() {
let bundle = output_only_bundle();
let output_count = bundle.outputs.len();
let parsed = bundle.into_parsed(AnchorRequirement::NotRequired).unwrap();
assert!(parsed.bundle.spends().is_empty());
assert_eq!(parsed.bundle.outputs().len(), output_count);
let reserialized = parsed.reserialize();
assert!(reserialized.anchor.is_none());
assert_eq!(reserialized.outputs.len(), output_count);
}
}