use crate::{KIND_ADD_FULL, KIND_ADD_HI, KIND_BASIC, KIND_SH3ADD_ADD, KIND_SH3ADD_HI};
use zisk_core::zisk_ops::ZiskOp;
pub const NEG_HI: u64 = 0xFFFF_FFFF;
const MASK_32: u64 = 0xFFFF_FFFF;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum AddShape {
Hi,
HiNeg,
Full,
}
#[inline(always)]
pub fn add_shape(a: u64, b: u64) -> AddShape {
let carry = ((a & MASK_32) + (b & MASK_32)) >> 32;
if (a >> 32) != 0 {
return AddShape::Full;
}
match b >> 32 {
0 if carry == 0 => AddShape::Hi,
NEG_HI if carry == 1 => AddShape::HiNeg,
_ => AddShape::Full,
}
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum Sh3addShape {
Hi,
Add,
Full,
}
#[inline(always)]
pub fn sh3add_shape(a: u64, b: u64) -> Sh3addShape {
if (a >> 32) != 0 {
return Sh3addShape::Full;
}
let sum = ((a & MASK_32) << 3) + (b & MASK_32);
let carry = sum >> 32;
match b >> 32 {
0 if carry == 0 => Sh3addShape::Hi,
NEG_HI if carry == 1 => Sh3addShape::Hi,
_ if carry <= 1 => Sh3addShape::Add,
_ => Sh3addShape::Full,
}
}
#[inline(always)]
pub fn add_family_kind(op: u8, a: u64, b: u64) -> usize {
if op == ZiskOp::Add.code() {
match add_shape(a, b) {
AddShape::Hi | AddShape::HiNeg => KIND_ADD_HI,
AddShape::Full => KIND_ADD_FULL,
}
} else if op == ZiskOp::Sh3add.code() {
match sh3add_shape(a, b) {
Sh3addShape::Hi => KIND_SH3ADD_HI,
Sh3addShape::Add => KIND_SH3ADD_ADD,
Sh3addShape::Full => KIND_BASIC,
}
} else {
KIND_BASIC
}
}
pub fn opcode_is_shift(opcode: ZiskOp) -> bool {
match opcode {
ZiskOp::Sll
| ZiskOp::Srl
| ZiskOp::Sra
| ZiskOp::SllW
| ZiskOp::SrlW
| ZiskOp::SraW
| ZiskOp::Rol
| ZiskOp::RolW
| ZiskOp::Ror
| ZiskOp::RorW
| ZiskOp::Bclr
| ZiskOp::Bext
| ZiskOp::Binv
| ZiskOp::Bset
| ZiskOp::SllUW => true,
ZiskOp::SignExtendB
| ZiskOp::SignExtendH
| ZiskOp::SignExtendW
| ZiskOp::Rev8
| ZiskOp::OrcB
| ZiskOp::Cpop
| ZiskOp::CpopW
| ZiskOp::Ctz
| ZiskOp::CtzW
| ZiskOp::Clz
| ZiskOp::ClzW
| ZiskOp::Pack
| ZiskOp::PackH
| ZiskOp::PackW => false,
_ => panic!("opcode_is_shift() got invalid opcode={opcode:?}"),
}
}
pub fn opcode_is_chain(opcode: ZiskOp) -> bool {
matches!(opcode, ZiskOp::Ctz | ZiskOp::CtzW)
}
pub fn opcode_is_chain_rev(opcode: ZiskOp) -> bool {
matches!(opcode, ZiskOp::Clz | ZiskOp::ClzW)
}
pub fn opcode_is_combine(opcode: ZiskOp) -> bool {
matches!(opcode, ZiskOp::Pack | ZiskOp::PackH | ZiskOp::PackW)
}
pub fn opcode_is_shift_word(opcode: ZiskOp) -> bool {
match opcode {
ZiskOp::SllW | ZiskOp::SrlW | ZiskOp::SraW | ZiskOp::RolW | ZiskOp::RorW => true,
ZiskOp::Sll
| ZiskOp::Srl
| ZiskOp::Sra
| ZiskOp::SllUW
| ZiskOp::SignExtendB
| ZiskOp::SignExtendH
| ZiskOp::SignExtendW
| ZiskOp::Rev8
| ZiskOp::OrcB
| ZiskOp::Rol
| ZiskOp::Ror
| ZiskOp::Cpop
| ZiskOp::CpopW
| ZiskOp::Ctz
| ZiskOp::CtzW
| ZiskOp::Clz
| ZiskOp::ClzW
| ZiskOp::Pack
| ZiskOp::PackH
| ZiskOp::PackW
| ZiskOp::Bclr
| ZiskOp::Bext
| ZiskOp::Binv
| ZiskOp::Bset => false,
_ => panic!("opcode_is_shift_word() got invalid opcode={opcode:?}"),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sh3add_reference(a: u64, b: u64) -> Sh3addShape {
let c = b.wrapping_add(a.wrapping_shl(3));
let a_is_clean = (a >> 32) == 0;
let hi_provable = a_is_clean && (c >> 32) == 0 && matches!(b >> 32, 0 | NEG_HI);
let add_provable = a_is_clean && ((a & MASK_32) << 3) + (b & MASK_32) < (1u64 << 33);
if hi_provable {
Sh3addShape::Hi
} else if add_provable {
Sh3addShape::Add
} else {
Sh3addShape::Full
}
}
#[test]
fn sh3add_shape_matches_what_the_airs_prove() {
let interesting = [
0u64,
1,
7,
8,
0x1FFF_FFFF, 0x2000_0000,
0xFFFF_FFFF,
0x1_0000_0000, u64::MAX, 0xFFFF_FFFF_0000_0008,
0xA000_0000, 0xFFFF_FFFE,
];
for &a in &interesting {
for &b in &interesting {
assert_eq!(sh3add_shape(a, b), sh3add_reference(a, b), "a=0x{a:X} b=0x{b:X}");
}
}
}
#[test]
fn sh3add_shape_matches_on_random_operands() {
let mut state = 0x2545_F491_4F6C_DD1Du64;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for _ in 0..200_000 {
let r = next();
let a = if r & 3 == 0 { r } else { r & MASK_32 };
let b = if r & 12 == 0 { next() } else { next() & MASK_32 };
assert_eq!(sh3add_shape(a, b), sh3add_reference(a, b), "a=0x{a:X} b=0x{b:X}");
}
}
#[test]
fn sh3add_address_arithmetic_is_the_cheap_shape() {
let base = 0xA000_1000u64;
for index in [0u64, 1, 2, 100, 1000, 0x10_0000] {
assert_eq!(sh3add_shape(index, base), Sh3addShape::Hi, "index={index}");
}
}
#[test]
fn sh3add_a_must_be_a_clean_32_bit_value() {
assert_eq!(sh3add_shape(1 << 32, 0), Sh3addShape::Full);
assert_eq!(sh3add_shape(u64::MAX, 0), Sh3addShape::Full, "negative a is dirty too");
}
#[test]
fn every_kind_is_provable_by_the_airs_that_claim_it() {
use crate::{add_family, KIND_BASIC};
use zisk_core::zisk_ops::ZiskOp;
let interesting = [
0u64,
1,
8,
0x1FFF_FFFF,
0x2000_0000,
0xA000_1000,
0xFFFF_FFFF,
0x1_0000_0000,
u64::MAX,
0xFFFF_FFFF_0000_0008,
];
let airs = add_family([1; crate::ADD_AIRS]);
for &op in &[ZiskOp::Add.code(), ZiskOp::Sh3add.code()] {
for &a in &interesting {
for &b in &interesting {
let kind = add_family_kind(op, a, b);
assert!(
airs.iter().any(|air| air.proves[kind]),
"op={op:#x} a={a:#X} b={b:#X} lands on kind {kind}, which no air proves",
);
if kind != KIND_BASIC {
assert!(
airs.iter().filter(|air| !air.proves[KIND_BASIC]).any(|a| a.proves[kind]),
"op={op:#x} a={a:#X} b={b:#X} is kind {kind} but no packed air proves it",
);
}
}
}
}
}
#[test]
fn other_opcodes_are_always_basic() {
use zisk_core::zisk_ops::ZiskOp;
for op in
[ZiskOp::And, ZiskOp::Or, ZiskOp::Xor, ZiskOp::Sub, ZiskOp::Sh1add, ZiskOp::Sh2add]
{
assert_eq!(add_family_kind(op.code(), 1, 2), crate::KIND_BASIC, "{op:?}");
}
}
#[test]
fn add_shape_hi() {
assert_eq!(add_shape(0, 0), AddShape::Hi);
assert_eq!(add_shape(1, 2), AddShape::Hi);
assert_eq!(add_shape(0xFFFF_FFFE, 1), AddShape::Hi);
}
#[test]
fn add_shape_hi_carry_out_of_low_limb_is_not_hi() {
assert_eq!(add_shape(0xFFFF_FFFF, 1), AddShape::Full);
}
#[test]
fn add_shape_hi_neg() {
let minus_one = u64::MAX;
assert_eq!(add_shape(1, minus_one), AddShape::HiNeg);
assert_eq!(add_shape(0xFFFF_FFFF, minus_one), AddShape::HiNeg);
assert_eq!(add_shape(0, minus_one), AddShape::Full);
}
#[test]
fn add_shape_full_when_a_is_dirty() {
assert_eq!(add_shape(1 << 32, 0), AddShape::Full);
}
#[test]
fn add_shape_full_when_b_hi_is_neither_zero_nor_all_ones() {
assert_eq!(add_shape(0, 1 << 32), AddShape::Full);
}
}