use alloc::{collections::BTreeMap, vec::Vec};
use core::array;
use miden_core::{
Felt,
field::{PrimeCharacteristicRing, QuadFelt},
utils::{Matrix, RowMajorMatrix},
};
use super::{
AUX_WIDTH, CARRY_HI_BEGIN, CARRY_LO_BEGIN, HUB_CELL_UINTLIMBS_MULT, HUB_CELL_UINTVAL_MULT,
NUM_CELLS, NUM_MAIN_COLS, PERIOD, TERM_CELL_GAP, UintStoreAir,
};
use crate::{
logup::build_logup_aux_trace,
math::{U256, to_limbs16, to_limbs32},
primitives::byte_pair_lut::BytePairLutRequires,
relations::ProvideMult,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct UintPtr(u32);
impl UintPtr {
pub fn addr(self) -> u32 {
self.0
}
pub fn from_addr(addr: u32) -> Self {
Self(addr)
}
}
#[derive(Debug, Clone, Copy)]
pub struct Uint {
pub value: U256,
pub ptr: UintPtr,
pub bound_ptr: UintPtr,
}
fn carries(v32: &[u32; 8], comp32: &[u32; 8]) -> [u16; 7] {
let mut c = [0u16; 7];
let mut carry: u64 = 0;
for j in 0..7 {
let s = v32[j] as u64 + comp32[j] as u64 + carry;
carry = s >> 32;
c[j] = carry as u16;
}
c
}
#[derive(Debug, Default)]
pub struct UintValRequires {
demand: BTreeMap<UintPtr, ProvideMult>,
}
impl UintValRequires {
pub fn new() -> Self {
Self::default()
}
pub fn require(&mut self, ptr: UintPtr) {
*self.demand.entry(ptr).or_insert(0) += 1;
}
pub fn count(&self, ptr: UintPtr) -> ProvideMult {
self.demand.get(&ptr).copied().unwrap_or(0)
}
}
pub const PIN_NAMESPACE_END: u32 = 1 << 16;
#[derive(Debug)]
pub struct UintStoreRequires {
uints: BTreeMap<UintPtr, Uint>,
by_value: BTreeMap<(U256, UintPtr), UintPtr>,
next_transient: u32,
demand: UintValRequires,
limbs_demand: UintValRequires,
}
impl UintStoreRequires {
pub fn new() -> Self {
Self {
uints: BTreeMap::new(),
by_value: BTreeMap::new(),
next_transient: PIN_NAMESPACE_END,
demand: UintValRequires::new(),
limbs_demand: UintValRequires::new(),
}
}
pub fn pin_modulus(&mut self, addr: u32, bound: U256) -> UintPtr {
let ptr = UintPtr(addr);
self.insert_pinned(ptr, bound, ptr, false);
ptr
}
pub fn intern_pinned(&mut self, addr: u32, value: U256, bound: UintPtr) -> UintPtr {
let ptr = UintPtr(addr);
assert!(value <= self.uint(bound).value, "value exceeds its modulus bound");
self.insert_pinned(ptr, value, bound, false);
ptr
}
pub fn intern_fixed_pinned(&mut self, addr: u32, value: U256, bound: UintPtr) -> UintPtr {
let ptr = UintPtr(addr);
assert!(value <= self.uint(bound).value, "value exceeds its modulus bound");
self.insert_pinned(ptr, value, bound, true);
ptr
}
fn insert_pinned(
&mut self,
ptr: UintPtr,
value: U256,
bound_ptr: UintPtr,
allow_value_alias: bool,
) {
assert!(
(1..PIN_NAMESPACE_END).contains(&ptr.0),
"pinned uint ptr {} outside the pin namespace [1, 2^16)",
ptr.0,
);
assert!(!self.uints.contains_key(&ptr), "duplicate uint ptr {}", ptr.0,);
match self.by_value.get(&(value, bound_ptr)).copied() {
Some(prev) if allow_value_alias => {
debug_assert_ne!(prev, ptr);
},
Some(prev) => {
panic!("value already interned at ptr {} — pin before computing", prev.0);
},
None => {
self.by_value.insert((value, bound_ptr), ptr);
},
}
self.uints.insert(ptr, Uint { value, ptr, bound_ptr });
self.demand.require(bound_ptr);
}
pub fn pinned(&self, addr: u32) -> UintPtr {
assert!(
(1..PIN_NAMESPACE_END).contains(&addr),
"address {addr} outside the pin namespace [1, 2^16)",
);
let ptr = UintPtr(addr);
assert!(self.uints.contains_key(&ptr), "no uint pinned at address {addr}",);
ptr
}
pub fn require_uintval(&mut self, ptr: UintPtr) {
self.demand.require(ptr);
}
pub fn require_uintlimbs(&mut self, ptr: UintPtr) {
self.limbs_demand.require(ptr);
}
pub fn intern(&mut self, value: U256, bound: UintPtr) -> UintPtr {
if let Some(&ptr) = self.by_value.get(&(value, bound)) {
return ptr;
}
assert!(value <= self.uint(bound).value, "value exceeds its modulus bound");
let ptr = UintPtr(self.next_transient);
self.next_transient += 1;
self.by_value.insert((value, bound), ptr);
self.uints.insert(ptr, Uint { value, ptr, bound_ptr: bound });
self.demand.require(bound);
ptr
}
pub fn uint(&self, ptr: UintPtr) -> &Uint {
self.uints
.get(&ptr)
.unwrap_or_else(|| panic!("ptr {} is not an interned uint", ptr.0))
}
}
impl Default for UintStoreRequires {
fn default() -> Self {
Self::new()
}
}
fn padded_blocks(requires: &UintStoreRequires, min_blocks: usize) -> Vec<Uint> {
let n_real = requires.uints.len();
let n_padded = n_real.next_power_of_two().max(1).max(min_blocks);
let next_ptr = requires.uints.last_key_value().map_or(1, |(&ptr, _)| ptr.0 + 1);
let pad = (0..n_padded - n_real).map(|i| {
let ptr = UintPtr(next_ptr + i as u32);
Uint { value: U256::ZERO, ptr, bound_ptr: ptr }
});
requires.uints.values().copied().chain(pad).collect()
}
fn bound_value(requires: &UintStoreRequires, u: &Uint, is_pad: bool) -> U256 {
if is_pad {
return U256::ZERO;
}
requires.uint(u.bound_ptr).value
}
pub fn generate_trace(
requires: UintStoreRequires,
bpl: &mut BytePairLutRequires,
) -> RowMajorMatrix<Felt> {
generate_trace_padded_to(requires, bpl, 0)
}
pub(crate) fn generate_trace_padded_to(
requires: UintStoreRequires,
bpl: &mut BytePairLutRequires,
min_blocks: usize,
) -> RowMajorMatrix<Felt> {
let requires = &requires;
let n_real = requires.uints.len();
let blocks = padded_blocks(requires, min_blocks);
let demand = &requires.demand;
let mut vals = Vec::with_capacity(blocks.len() * PERIOD * NUM_MAIN_COLS);
for (i, u) in blocks.iter().enumerate() {
let bound_value = bound_value(requires, u, i >= n_real);
let comp = bound_value.checked_sub(u.value).expect("stored value exceeds its bound");
let v16 = to_limbs16(u.value);
let comp16 = to_limbs16(comp);
let bound32 = to_limbs32(bound_value);
let c = carries(&to_limbs32(u.value), &to_limbs32(comp));
let mult = demand.count(u.ptr) + u32::from(i >= n_real);
let limbs_mult = requires.limbs_demand.count(u.ptr);
let gap = match blocks.get(i + 1) {
Some(nxt) => nxt.ptr.0 - u.ptr.0 - 1,
None => 0,
};
for l in v16.into_iter().chain(comp16) {
bpl.require_range16(l);
}
bpl.require_range16(gap as u16);
let mut v_lo = [Felt::ZERO; NUM_CELLS];
for i in 0..8 {
v_lo[i] = Felt::from(v16[i]);
}
let mut v_hi = [Felt::ZERO; NUM_CELLS];
for i in 0..8 {
v_hi[i] = Felt::from(v16[8 + i]);
}
v_hi[HUB_CELL_UINTVAL_MULT] = Felt::from(mult);
v_hi[HUB_CELL_UINTLIMBS_MULT] = Felt::from(limbs_mult);
let comp: [Felt; NUM_CELLS] = array::from_fn(|i| Felt::from(comp16[i]));
let mut bound = [Felt::ZERO; NUM_CELLS];
for i in 0..4 {
bound[i] = Felt::from(bound32[i]);
bound[CARRY_LO_BEGIN + i] = Felt::from(c[i]);
bound[8 + i] = Felt::from(bound32[4 + i]);
}
for j in 0..3 {
bound[CARRY_HI_BEGIN + j] = Felt::from(c[4 + j]);
}
bound[TERM_CELL_GAP] = Felt::from(gap);
let rows: [[Felt; NUM_CELLS]; PERIOD] = [v_lo, v_hi, comp, bound];
for row in rows {
vals.extend(row);
vals.push(Felt::from(u.ptr.0));
vals.push(Felt::from(u.bound_ptr.0));
}
}
RowMajorMatrix::new(vals, NUM_MAIN_COLS)
}
pub(crate) fn build_aux(
main: &RowMajorMatrix<Felt>,
challenges: &[QuadFelt],
) -> (RowMajorMatrix<QuadFelt>, Vec<QuadFelt>) {
let (logup, sigma) = build_logup_aux_trace(&UintStoreAir, main, challenges);
let n = main.height();
let beta = challenges[1];
let mut bp = [QuadFelt::ZERO; 8];
bp[0] = QuadFelt::ONE;
for i in 1..8 {
bp[i] = bp[i - 1] * beta;
}
let two16 = Felt::from(1u32 << 16);
let t32 = QuadFelt::from(Felt::new(1u64 << 32).expect("2^32 < Goldilocks p"));
let logup_width = logup.width();
let mut data = Vec::with_capacity(AUX_WIDTH * n);
let mut id = QuadFelt::ZERO;
for r in 0..n {
data.extend((0..logup_width).map(|c| logup.values[r * logup_width + c]));
data.push(id);
let limb = |c: usize| -> Felt { main.values[r * NUM_MAIN_COLS + c] };
let recomb_lo07 = || {
(0..4).fold(QuadFelt::ZERO, |s, k| {
let rk = limb(2 * k) + two16 * limb(2 * k + 1);
s + bp[k] * QuadFelt::from(rk)
})
};
let recomb_hi07 = || {
(0..4).fold(QuadFelt::ZERO, |s, k| {
let rk = limb(2 * k) + two16 * limb(2 * k + 1);
s + bp[4 + k] * QuadFelt::from(rk)
})
};
let recomb_hi815 = || {
(0..4).fold(QuadFelt::ZERO, |s, k| {
let rk = limb(8 + 2 * k) + two16 * limb(8 + 2 * k + 1);
s + bp[4 + k] * QuadFelt::from(rk)
})
};
let contrib: QuadFelt = match r % PERIOD {
0 => recomb_lo07(),
1 => recomb_hi07(),
2 => recomb_lo07() + recomb_hi815(),
3 => {
let carry_lo = (0..4).fold(QuadFelt::ZERO, |s, j| {
let w = bp[j + 1] - bp[j] * t32;
s + w * QuadFelt::from(limb(CARRY_LO_BEGIN + j))
});
let carry_hi = (0..3).fold(QuadFelt::ZERO, |s, j| {
let w = bp[4 + j + 1] - bp[4 + j] * t32;
s + w * QuadFelt::from(limb(CARRY_HI_BEGIN + j))
});
let direct_lo =
(0..4).fold(QuadFelt::ZERO, |s, k| s + bp[k] * QuadFelt::from(limb(k)));
let direct_hi =
(0..4).fold(QuadFelt::ZERO, |s, k| s + bp[4 + k] * QuadFelt::from(limb(8 + k)));
carry_lo - direct_lo + carry_hi - direct_hi
},
_ => unreachable!("PERIOD = 4"),
};
id += contrib;
}
(RowMajorMatrix::new(data, AUX_WIDTH), sigma)
}