use core::ops::Deref as _;
use std_shims::{vec, vec::Vec, collections::HashMap};
use zeroize::{Zeroize, ZeroizeOnDrop, Zeroizing};
#[cfg(feature = "compile-time-generators")]
use curve25519_dalek::constants::ED25519_BASEPOINT_TABLE;
#[cfg(not(feature = "compile-time-generators"))]
use curve25519_dalek::constants::ED25519_BASEPOINT_POINT as ED25519_BASEPOINT_TABLE;
use monero_oxide::{
ed25519::{Scalar, CompressedPoint, Point, Commitment},
transaction::{Timelock, Pruned, Transaction},
};
use monero_interface::ScannableBlock;
use crate::{
address::SubaddressIndex, ViewPair, GuaranteedViewPair, output::*, PaymentId, Extra,
SharedKeyDerivations,
};
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct Timelocked(Vec<WalletOutput>);
impl Timelocked {
#[must_use]
pub fn not_additionally_locked(self) -> Vec<WalletOutput> {
let mut res = vec![];
for output in &self.0 {
if output.additional_timelock() == Timelock::None {
res.push(output.clone());
}
}
res
}
#[must_use]
pub fn additional_timelock_satisfied_by(self, block: usize, time: u64) -> Vec<WalletOutput> {
let mut res = vec![];
for output in &self.0 {
if (output.additional_timelock() <= Timelock::Block(block)) ||
(output.additional_timelock() <= Timelock::Time(time))
{
res.push(output.clone());
}
}
res
}
#[must_use]
pub fn ignore_additional_timelock(mut self) -> Vec<WalletOutput> {
let mut res = vec![];
core::mem::swap(&mut self.0, &mut res);
res
}
}
#[derive(Clone, Copy, PartialEq, Eq, Debug, thiserror::Error)]
pub enum ScanError {
#[error("unsupported protocol version ({0})")]
UnsupportedProtocol(u8),
#[error("invalid scannable block ({0})")]
InvalidScannableBlock(&'static str),
}
#[derive(Clone)]
struct InternalScanner {
pair: ViewPair,
guaranteed: bool,
subaddresses: HashMap<CompressedPoint, Option<SubaddressIndex>>,
}
impl Zeroize for InternalScanner {
#[expect(clippy::iter_over_hash_type)]
fn zeroize(&mut self) {
self.pair.zeroize();
self.guaranteed.zeroize();
for (mut key, mut value) in self.subaddresses.drain() {
key.zeroize();
value.zeroize();
}
}
}
impl Drop for InternalScanner {
fn drop(&mut self) {
self.zeroize();
}
}
impl ZeroizeOnDrop for InternalScanner {}
impl InternalScanner {
fn new(pair: ViewPair, guaranteed: bool) -> Self {
let mut subaddresses = HashMap::new();
subaddresses.insert(pair.spend().compress(), None);
Self { pair, guaranteed, subaddresses }
}
fn register_subaddress(&mut self, subaddress: SubaddressIndex) {
let (spend, _) = self.pair.subaddress_keys(subaddress);
self.subaddresses.insert(spend.compress(), Some(subaddress));
}
fn scan_transaction(
&self,
output_index_for_first_ringct_output: u64,
tx_hash: [u8; 32],
tx: &Transaction<Pruned>,
) -> Result<Timelocked, ScanError> {
if tx.version() != 2 {
return Ok(Timelocked(vec![]));
}
let Ok(extra) = Extra::read(&mut tx.prefix().extra.as_slice()) else {
return Ok(Timelocked(vec![]));
};
let Some((tx_keys, additional)) = extra.keys() else {
return Ok(Timelocked(vec![]));
};
let payment_id = extra.payment_id();
let mut res = vec![];
for (o, output) in tx.prefix().outputs.iter().enumerate() {
if output.key == CompressedPoint::IDENTITY {
continue;
}
let Some(output_key) = output.key.decompress() else { continue };
let additional = additional.as_ref().and_then(|additional| additional.get(o));
for key in tx_keys.iter().map(Some).chain(core::iter::once(additional)).flatten().copied() {
let ecdh = {
let dalek_view = Zeroizing::new((*self.pair.view).into());
Zeroizing::new(Point::from(dalek_view.deref() * key.into()))
};
let output_derivations = SharedKeyDerivations::output_derivations(
self.guaranteed.then(|| SharedKeyDerivations::uniqueness(&tx.prefix().inputs)),
ecdh.clone(),
o,
);
if let Some(actual_view_tag) = output.view_tag {
if actual_view_tag != output_derivations.view_tag {
continue;
}
}
let Some(subaddress) = ({
let subaddress_spend_key =
output_key.into() - (&output_derivations.shared_key.into() * ED25519_BASEPOINT_TABLE);
self
.subaddresses
.get::<CompressedPoint>(&subaddress_spend_key.compress().to_bytes().into())
}) else {
continue;
};
let subaddress = *subaddress;
let mut key_offset = output_derivations.shared_key.into();
if let Some(subaddress) = subaddress {
key_offset += self.pair.subaddress_derivation(subaddress).into();
}
let mut commitment = Commitment::zero();
if let Some(amount) = output.amount {
commitment.amount = amount;
} else {
let Transaction::V2 { proofs: Some(ref proofs), .. } = &tx else {
Err(ScanError::InvalidScannableBlock("non-miner v2 transaction without RCT proofs"))?
};
commitment = match proofs.base.encrypted_amounts.get(o) {
Some(amount) => output_derivations.decrypt(amount),
None => Err(ScanError::InvalidScannableBlock(
"RCT proofs without an encrypted amount per output",
))?,
};
if Some(&commitment.commit().compress()) != proofs.base.commitments.get(o) {
continue;
}
}
let payment_id = payment_id.map(|id| id ^ SharedKeyDerivations::payment_id_xor(ecdh));
let o = u64::try_from(o).expect("couldn't convert output index (usize) to u64");
res.push(WalletOutput {
absolute_id: AbsoluteId { transaction: tx_hash, index_in_transaction: o },
relative_id: RelativeId {
index_on_blockchain: output_index_for_first_ringct_output.checked_add(o).ok_or(
ScanError::InvalidScannableBlock(
"transaction's output's index isn't representable as a u64",
),
)?,
},
data: OutputData { key: output_key, key_offset: Scalar::from(key_offset), commitment },
metadata: Metadata {
additional_timelock: tx.prefix().additional_timelock,
subaddress,
payment_id,
arbitrary_data: extra.arbitrary_data(),
},
});
break;
}
}
Ok(Timelocked(res))
}
fn scan(&mut self, block: ScannableBlock) -> Result<Timelocked, ScanError> {
let ScannableBlock { block, transactions, output_index_for_first_ringct_output } = block;
if block.transactions.len() != transactions.len() {
Err(ScanError::InvalidScannableBlock(
"scanning a ScannableBlock with more/less transactions than it should have",
))?;
}
let Some(mut output_index_for_first_ringct_output) = output_index_for_first_ringct_output
else {
return Ok(Timelocked(vec![]));
};
if block.header.hardfork_version > 16 {
Err(ScanError::UnsupportedProtocol(block.header.hardfork_version))?;
}
let mut txs_with_hashes = vec![(
block.miner_transaction().hash(),
Transaction::<Pruned>::from(block.miner_transaction().clone()),
)];
for (hash, tx) in block.transactions.iter().zip(transactions) {
txs_with_hashes.push((*hash, tx));
}
let mut res = Timelocked(vec![]);
for (hash, tx) in txs_with_hashes {
{
let mut this_txs_outputs = vec![];
core::mem::swap(
&mut self.scan_transaction(output_index_for_first_ringct_output, hash, &tx)?.0,
&mut this_txs_outputs,
);
res.0.extend(this_txs_outputs);
}
if matches!(tx, Transaction::V2 { .. }) {
output_index_for_first_ringct_output = output_index_for_first_ringct_output
.checked_add(
u64::try_from(tx.prefix().outputs.len())
.expect("couldn't convert amount of outputs (usize) to u64"),
)
.ok_or(ScanError::InvalidScannableBlock("RingCT output indexes exceeded u64::MAX"))?;
}
}
if block.header.hardfork_version >= 12 {
for output in &mut res.0 {
if matches!(output.metadata.payment_id, Some(PaymentId::Unencrypted(_))) {
output.metadata.payment_id = None;
}
}
}
Ok(res)
}
}
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct Scanner(InternalScanner);
impl Scanner {
pub fn new(pair: ViewPair) -> Self {
Self(InternalScanner::new(pair, false))
}
pub fn register_subaddress(&mut self, subaddress: SubaddressIndex) {
self.0.register_subaddress(subaddress);
}
pub fn scan(&mut self, block: ScannableBlock) -> Result<Timelocked, ScanError> {
self.0.scan(block)
}
}
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct GuaranteedScanner(InternalScanner);
impl GuaranteedScanner {
pub fn new(pair: GuaranteedViewPair) -> Self {
Self(InternalScanner::new(pair.0, true))
}
pub fn register_subaddress(&mut self, subaddress: SubaddressIndex) {
self.0.register_subaddress(subaddress);
}
pub fn scan(&mut self, block: ScannableBlock) -> Result<Timelocked, ScanError> {
self.0.scan(block)
}
}