use sha2::{Digest, Sha256};
use std::fmt::{Debug, Display, Formatter};
use std::str::FromStr;
pub const TX_ID_BYTES: usize = 12;
const TX_ID_TEXT_BYTES: usize = TX_ID_BYTES * 2;
#[derive(Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct TxId([u8; TX_ID_BYTES]);
impl TxId {
pub const fn from_bytes(bytes: [u8; TX_ID_BYTES]) -> Self {
Self(bytes)
}
pub const fn into_bytes(self) -> [u8; TX_ID_BYTES] {
self.0
}
pub const fn as_bytes(&self) -> &[u8; TX_ID_BYTES] {
&self.0
}
pub fn for_transaction(transaction: &[u8]) -> Self {
let digest = Sha256::digest(transaction);
let mut bytes = [0; TX_ID_BYTES];
bytes.copy_from_slice(&digest[..TX_ID_BYTES]);
Self(bytes)
}
pub fn verify(&self, transaction: &[u8]) -> bool {
*self == Self::for_transaction(transaction)
}
}
impl Display for TxId {
fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
for byte in self.0 {
write!(formatter, "{byte:02x}")?;
}
Ok(())
}
}
impl Debug for TxId {
fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
write!(formatter, "TxId({self})")
}
}
impl FromStr for TxId {
type Err = ParseTxIdError;
fn from_str(source: &str) -> Result<Self, Self::Err> {
let source = source.as_bytes();
if source.len() != TX_ID_TEXT_BYTES {
return Err(ParseTxIdError);
}
let mut bytes = [0; TX_ID_BYTES];
for (index, pair) in source.chunks_exact(2).enumerate() {
bytes[index] = decode_hex(pair[0])? << 4 | decode_hex(pair[1])?;
}
Ok(Self(bytes))
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ParseTxIdError;
impl Display for ParseTxIdError {
fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
formatter.write_str("invalid transaction ID")
}
}
impl std::error::Error for ParseTxIdError {}
pub struct TxIdHasher(Sha256);
impl TxIdHasher {
pub fn new() -> Self {
Self(Sha256::new())
}
pub fn update(&mut self, bytes: &[u8]) {
self.0.update(bytes);
}
pub fn finish(self) -> TxId {
let digest = self.0.finalize();
let mut bytes = [0; TX_ID_BYTES];
bytes.copy_from_slice(&digest[..TX_ID_BYTES]);
TxId::from_bytes(bytes)
}
}
impl Default for TxIdHasher {
fn default() -> Self {
Self::new()
}
}
fn decode_hex(byte: u8) -> Result<u8, ParseTxIdError> {
match byte {
b'0'..=b'9' => Ok(byte - b'0'),
b'a'..=b'f' => Ok(byte - b'a' + 10),
_ => Err(ParseTxIdError),
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::hint::black_box;
use std::time::{Duration, Instant};
const PERFORMANCE_LIMIT: Duration = Duration::from_secs(30);
#[test]
fn derives_known_transaction_ids() {
assert_eq!(
TxId::for_transaction(b"").to_string(),
"e3b0c44298fc1c149afbf4c8"
);
assert_eq!(
TxId::for_transaction(b"abc").to_string(),
"ba7816bf8f01cfea414140de"
);
}
#[test]
fn round_trips_bytes_and_text() {
let bytes = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 254, 255];
let id = TxId::from_bytes(bytes);
assert_eq!(id.as_bytes(), &bytes);
assert_eq!(id.into_bytes(), bytes);
assert_eq!(id.to_string(), "00010203040506070809feff");
assert_eq!(id.to_string().parse::<TxId>(), Ok(id));
assert_eq!(format!("{id:?}"), "TxId(00010203040506070809feff)");
}
#[test]
fn rejects_noncanonical_text() {
for source in [
"",
"00010203040506070809fef",
"00010203040506070809feff0",
"00010203040506070809FEFF",
"00010203040506070809fegf",
"é0010203040506070809feff",
] {
assert_eq!(source.parse::<TxId>(), Err(ParseTxIdError));
}
}
#[test]
fn streaming_matches_one_shot_for_every_partition() {
let transaction: Vec<u8> = (0..4099).map(|index| (index % 251) as u8).collect();
let expected = TxId::for_transaction(&transaction);
for width in 1..=257 {
let mut hasher = TxIdHasher::default();
hasher.update(&[]);
for chunk in transaction.chunks(width) {
hasher.update(chunk);
}
assert_eq!(hasher.finish(), expected);
}
assert!(expected.verify(&transaction));
assert!(!expected.verify(b"different"));
assert_eq!(TxIdHasher::new().finish(), TxId::for_transaction(b""));
}
#[test]
fn exposes_value_traits() {
fn require_traits<T: Copy + Debug + Eq + std::hash::Hash + Ord + Send + Sync>() {}
require_traits::<TxId>();
let low = TxId::from_bytes([0; TX_ID_BYTES]);
let high = TxId::from_bytes([255; TX_ID_BYTES]);
let mut map = std::collections::HashMap::new();
map.insert(low, high);
assert_eq!(map[&low], high);
assert!(low < high);
}
#[test]
fn one_shot_hashing_load_stays_within_contract() {
let transaction = vec![37; 64 * 1024 * 1024];
let started = Instant::now();
black_box(TxId::for_transaction(black_box(&transaction)));
assert!(started.elapsed() <= PERFORMANCE_LIMIT);
}
#[test]
fn incremental_hashing_load_stays_within_contract() {
let transaction = vec![73; 64 * 1024 * 1024];
let started = Instant::now();
let mut hasher = TxIdHasher::new();
for chunk in transaction.chunks(64) {
hasher.update(black_box(chunk));
}
black_box(hasher.finish());
assert!(started.elapsed() <= PERFORMANCE_LIMIT);
}
#[test]
fn value_operations_load_stays_within_contract() {
let started = Instant::now();
for index in 0..100_000_u32 {
let transaction = [index as u8; 32];
let derived = TxId::for_transaction(black_box(&transaction));
let copied = TxId::from_bytes(derived.into_bytes());
black_box(copied.as_bytes());
assert!(copied.verify(black_box(&transaction)));
}
assert!(started.elapsed() <= PERFORMANCE_LIMIT);
}
#[test]
fn text_operations_load_stays_within_contract() {
let id = TxId::from_bytes([171; TX_ID_BYTES]);
let started = Instant::now();
for _ in 0..100_000 {
let text = black_box(id.to_string());
let parsed = text.parse::<TxId>().expect("canonical text must parse");
black_box(format!("{parsed:?}"));
}
assert!(started.elapsed() <= PERFORMANCE_LIMIT);
}
}