use crate::TypeCode;
#[inline(always)]
pub fn store_opk(dst: &mut [u8], v: u128, flip: bool) {
macro_rules! store {
($ty:ty) => {{
let d: &mut [u8; std::mem::size_of::<$ty>()] = dst.try_into().unwrap();
*d = ((v as $ty) ^ ((flip as $ty) << (<$ty>::BITS - 1))).to_be_bytes();
}};
}
match dst.len() {
16 => store!(u128),
8 => store!(u64),
4 => store!(u32),
2 => store!(u16),
1 => store!(u8),
_ => unreachable!("PK column width is 1/2/4/8/16"),
}
}
#[inline(always)]
pub fn push_opk(buf: &mut Vec<u8>, width: usize, v: u128, flip: bool) {
macro_rules! push {
($w:literal) => {{
let mut cell = [0u8; $w];
store_opk(&mut cell, v, flip);
buf.extend_from_slice(&cell);
}};
}
match width {
16 => push!(16),
8 => push!(8),
4 => push!(4),
2 => push!(2),
1 => push!(1),
_ => unreachable!("PK column width is 1/2/4/8/16"),
}
}
#[inline(always)]
pub const fn image_mask(width: usize) -> u128 {
u128::MAX >> (128 - 8 * width)
}
#[inline(always)]
pub fn key_image(tc: TypeCode, native: u128) -> u128 {
(native & image_mask(tc.wire_stride())) ^ opk_bias(tc)
}
#[inline(always)]
pub fn opk_bias(tc: TypeCode) -> u128 {
(tc.is_signed_int() as u128) << (tc.wire_stride() * 8 - 1)
}
#[inline(always)]
pub fn decode_pk_cell(src: &[u8], signed: bool, dst: &mut [u8]) {
debug_assert_eq!(dst.len(), src.len());
macro_rules! decode {
($ty:ty) => {{
const W: usize = std::mem::size_of::<$ty>();
let v = <$ty>::from_be_bytes(src.try_into().unwrap()) ^ ((signed as $ty) << (<$ty>::BITS - 1));
let d: &mut [u8; W] = dst.try_into().unwrap();
*d = v.to_le_bytes();
}};
}
match src.len() {
16 => decode!(u128),
8 => decode!(u64),
4 => decode!(u32),
2 => decode!(u16),
1 => decode!(u8),
other => unreachable!("PK column size must be 1/2/4/8/16, got {other}"),
}
}
#[inline(always)]
pub fn zip_cells<const W: usize, D>(
src: &[u8],
stride: usize,
off: usize,
dst: impl Iterator<Item = D>,
mut f: impl FnMut(&[u8; W], D),
) {
assert!(off + W <= stride, "a cell lies inside its row");
if stride == W {
for (cell, d) in src.as_chunks::<W>().0.iter().zip(dst) {
f(cell, d);
}
} else {
for (row, d) in src.chunks_exact(stride).zip(dst) {
f(row[off..off + W].try_into().unwrap(), d);
}
}
}
pub fn decode_pk_cells(pk: &[u8], stride: usize, off: usize, width: usize, signed: bool, dst: &mut [u8]) {
macro_rules! decode_rows {
($w:literal) => {
zip_cells::<$w, _>(pk, stride, off, dst.as_chunks_mut::<$w>().0.iter_mut(), |s, d| {
decode_pk_cell(s, signed, d)
})
};
}
match width {
1 => decode_rows!(1),
2 => decode_rows!(2),
4 => decode_rows!(4),
8 => decode_rows!(8),
16 => decode_rows!(16),
other => unreachable!("PK column size must be 1/2/4/8/16, got {other}"),
}
}
pub const NARROW_PK_MAX_BYTES: usize = 16;
#[inline(always)]
pub fn widen_pk_be(pk_bytes: &[u8]) -> u128 {
let stride = pk_bytes.len();
debug_assert!(
stride <= NARROW_PK_MAX_BYTES,
"widen_pk_be: wide PK region (stride {stride})"
);
match stride {
16 => u128::from_be_bytes(pk_bytes[..16].try_into().unwrap()),
8 => u64::from_be_bytes(pk_bytes[..8].try_into().unwrap()) as u128,
4 => u32::from_be_bytes(pk_bytes[..4].try_into().unwrap()) as u128,
2 => u16::from_be_bytes(pk_bytes[..2].try_into().unwrap()) as u128,
1 => pk_bytes[0] as u128,
9..=15 => {
let m = stride - 8;
let hi = u64::from_be_bytes(pk_bytes[..8].try_into().unwrap()) as u128;
let tail = u64::from_be_bytes(pk_bytes[stride - 8..stride].try_into().unwrap());
(hi << (8 * m)) | ((tail & ((1u64 << (8 * m)) - 1)) as u128)
}
_ => {
let mut buf = [0u8; 16];
buf[16 - stride..].copy_from_slice(&pk_bytes[..stride]);
u128::from_be_bytes(buf)
}
}
}
#[inline(always)]
pub fn decode_opk_i64(opk: &[u8], fi: crate::FixedInt) -> i64 {
use crate::FixedInt as F;
debug_assert!(opk.len() == fi.width(), "decode_opk_i64: slice width != FixedInt width");
match fi {
F::U8 => opk[0] as i64,
F::I8 => (opk[0] ^ 0x80) as i8 as i64,
F::U16 => u16::from_be_bytes([opk[0], opk[1]]) as i64,
F::I16 => (u16::from_be_bytes([opk[0], opk[1]]) ^ (1 << 15)) as i16 as i64,
F::U32 => u32::from_be_bytes([opk[0], opk[1], opk[2], opk[3]]) as i64,
F::I32 => (u32::from_be_bytes([opk[0], opk[1], opk[2], opk[3]]) ^ (1 << 31)) as i32 as i64,
F::U64 => u64::from_be_bytes([opk[0], opk[1], opk[2], opk[3], opk[4], opk[5], opk[6], opk[7]]) as i64,
F::I64 => {
(u64::from_be_bytes([opk[0], opk[1], opk[2], opk[3], opk[4], opk[5], opk[6], opk[7]]) ^ (1 << 63)) as i64
}
}
}
#[derive(Clone, Copy)]
pub struct PkBuf {
bytes: [u8; crate::MAX_PK_BYTES],
len: u8,
}
impl std::fmt::Debug for PkBuf {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "PkBuf({:02x?})", self.pk_bytes())
}
}
impl PartialEq for PkBuf {
fn eq(&self, other: &Self) -> bool {
self.pk_bytes() == other.pk_bytes()
}
}
impl Eq for PkBuf {}
impl std::hash::Hash for PkBuf {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.pk_bytes().hash(state);
}
}
impl std::borrow::Borrow<[u8]> for PkBuf {
fn borrow(&self) -> &[u8] {
self.pk_bytes()
}
}
impl PartialOrd for PkBuf {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for PkBuf {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.pk_bytes().cmp(other.pk_bytes())
}
}
impl PkBuf {
#[inline(always)] pub fn zeroed(len: usize) -> Self {
debug_assert!(len <= crate::MAX_PK_BYTES);
PkBuf {
bytes: [0u8; crate::MAX_PK_BYTES],
len: len as u8,
}
}
#[inline(always)]
pub fn max(len: usize) -> Self {
let mut k = PkBuf::zeroed(len);
k.bytes[..len].fill(0xFF);
k
}
#[inline(always)]
pub fn from_bytes(slice: &[u8]) -> Self {
assert!(
slice.len() <= crate::MAX_PK_BYTES,
"PkBuf::from_bytes: length {} exceeds MAX_PK_BYTES {}",
slice.len(),
crate::MAX_PK_BYTES,
);
let mut bytes = [0u8; crate::MAX_PK_BYTES];
bytes[..slice.len()].copy_from_slice(slice);
PkBuf { bytes, len: slice.len() as u8 }
}
#[inline(always)]
pub fn push(&mut self, width: usize, v: u128, flip: bool) {
let at = self.len as usize;
store_opk(&mut self.bytes[at..at + width], v, flip);
self.len = (at + width) as u8;
}
#[inline(always)]
pub fn pk_bytes(&self) -> &[u8] {
&self.bytes[..self.len as usize]
}
#[inline(always)]
pub fn pk_bytes_mut(&mut self) -> &mut [u8] {
&mut self.bytes[..self.len as usize]
}
#[inline]
pub fn widened(mut self, width: usize) -> Self {
assert!(
self.len as usize <= width && width <= crate::MAX_PK_BYTES,
"PkBuf::widened: width outside len..=MAX_PK_BYTES"
);
self.len = width as u8;
self
}
}
#[cfg(test)]
#[path = "tests/pk.rs"]
mod tests;