use std::sync::OnceLock;
use nebu::Goldilocks;
use hemera::constants::ROUND_CONSTANTS;
use hemera::field::{Goldilocks as HGold, MATRIX_DIAG_16};
use crate::types::{CCSInstance, CCSWitness, SparseMatrix};
pub const NUM_PATTERNS: usize = 18;
pub const NUM_ROUNDS: usize = 25;
pub const PARTIAL_FIRST: usize = 3;
pub const NUM_PARTIAL: usize = 16;
pub const Z_LEN: usize = 96;
pub const NUM_ROWS: usize = 64;
pub const fn reg_t(r: usize) -> usize {
r
}
pub const fn reg_t1(r: usize) -> usize {
r + 16
}
pub const CONST_IDX: usize = 32;
pub const fn sel(p: usize) -> usize {
33 + p
}
pub const fn round_sel(k: usize) -> usize {
51 + k
}
pub const IDX_PI: usize = 76;
pub const IDX_KAPPA: usize = 77;
pub const IDX_RC: usize = 78;
pub const fn cap_k(i: usize) -> usize {
79 + i
}
pub const fn cap_k1(i: usize) -> usize {
87 + i
}
pub const IDX_Y: usize = 95;
pub const fn sk(i: usize) -> usize {
match i {
0..=3 => reg_t(4 + i),
4..=7 => reg_t(10 + i - 4),
_ => cap_k(i - 8),
}
}
pub const fn sk1(i: usize) -> usize {
match i {
0..=3 => reg_t1(4 + i),
4..=7 => reg_t1(10 + i - 4),
_ => cap_k1(i - 8),
}
}
const ROW_SEL: usize = 0; const ROW_SEL_SUM: usize = 18;
const ROW_ROUND: usize = 19; const ROW_ROUND_SUM: usize = 44;
const ROW_PI: usize = 45;
const ROW_KAPPA: usize = 46;
const ROW_RC: usize = 47;
const ROW_PAT: usize = 48;
const A0: usize = 0;
const B0: usize = 1;
const C0: usize = 2;
const D0: usize = 3;
const E0: usize = 4;
const A1: usize = 5;
const B1: usize = 6;
const C1: usize = 7;
const D1: usize = 8;
const E1: usize = 9;
const F: usize = 10;
const G: usize = 11;
const H: usize = 12;
const I: usize = 13;
const P: usize = 14;
const NUM_MATRICES: usize = 15;
const CDE0: (usize, usize, usize) = (C0, D0, E0);
const CDE1: (usize, usize, usize) = (C1, D1, E1);
fn hg(h: HGold) -> Goldilocks {
Goldilocks::new(h.as_canonical_u64())
}
fn neg(x: Goldilocks) -> Goldilocks {
Goldilocks::ZERO - x
}
pub fn partial_rc(pr: usize) -> Goldilocks {
hg(ROUND_CONSTANTS[128 + pr])
}
type Form<'a> = &'a [(usize, Goldilocks)];
struct Builder {
m: Vec<SparseMatrix>,
}
impl Builder {
fn new() -> Self {
Self { m: (0..NUM_MATRICES).map(|_| SparseMatrix::new(NUM_ROWS, Z_LEN)).collect() }
}
fn put(&mut self, slot: usize, row: usize, form: Form) {
debug_assert!(
self.m[slot].entries[row].is_empty(),
"slot {slot} already placed on row {row}: two gadgets collide"
);
for &(col, c) in form {
self.m[slot].set(row, col, c);
}
}
fn lin(&mut self, row: usize, a: usize, b: usize, gate: usize, form: Form) {
self.put(a, row, &[(gate, Goldilocks::ONE)]);
self.put(b, row, form);
}
fn prod(&mut self, row: usize, (c, d, e): (usize, usize, usize), gate: usize, l: Form, r: Form) {
self.put(c, row, &[(gate, Goldilocks::ONE)]);
self.put(d, row, l);
self.put(e, row, r);
}
}
fn one() -> Goldilocks {
Goldilocks::ONE
}
fn m1() -> Goldilocks {
neg(Goldilocks::ONE)
}
fn build() -> CCSInstance {
let mut b = Builder::new();
for p in 0..NUM_PATTERNS {
b.lin(
ROW_SEL + p,
A0,
B0,
sel(p),
&[(reg_t(0), one()), (CONST_IDX, neg(Goldilocks::new(p as u64)))],
);
}
let mut sum: Vec<(usize, Goldilocks)> = (0..NUM_PATTERNS).map(|p| (sel(p), one())).collect();
sum.push((CONST_IDX, m1()));
b.put(P, ROW_SEL_SUM, &sum);
for k in 0..NUM_ROUNDS {
b.lin(
ROW_ROUND + k,
A0,
B0,
round_sel(k),
&[(reg_t(14), one()), (CONST_IDX, neg(Goldilocks::new(k as u64)))],
);
}
let mut usum: Vec<(usize, Goldilocks)> = (0..NUM_ROUNDS).map(|k| (round_sel(k), one())).collect();
usum.push((sel(15), m1()));
b.put(P, ROW_ROUND_SUM, &usum);
let mut pi: Vec<(usize, Goldilocks)> = vec![(IDX_PI, one())];
let mut rc: Vec<(usize, Goldilocks)> = vec![(IDX_RC, one())];
for pr in 0..NUM_PARTIAL {
pi.push((round_sel(PARTIAL_FIRST + pr), m1()));
rc.push((round_sel(PARTIAL_FIRST + pr), neg(partial_rc(pr))));
}
b.put(P, ROW_PI, &pi);
b.put(P, ROW_KAPPA, &[(IDX_KAPPA, one()), (sel(15), m1()), (round_sel(24), one())]);
b.put(P, ROW_RC, &rc);
b.prod(ROW_PAT, CDE0, IDX_PI, &[(IDX_Y, one())], &[(sk(0), one()), (IDX_RC, one())]);
b.lin(ROW_PAT, A0, B0, IDX_PI, &[(CONST_IDX, m1())]);
let diag: Vec<Goldilocks> = MATRIX_DIAG_16.iter().map(|&d| hg(d)).collect();
b.lin(
ROW_PAT + 1,
A0,
B0,
IDX_PI,
&[(sk1(0), one()), (IDX_Y, neg(diag[0])), (sk1(1), m1()), (sk(1), diag[1])],
);
for i in 2..16 {
b.lin(
ROW_PAT + i,
A0,
B0,
IDX_PI,
&[(sk1(i), one()), (sk(i), neg(diag[i])), (sk1(1), m1()), (sk(1), diag[1])],
);
}
b.lin(ROW_PAT, A1, B1, sel(0), &[(reg_t(9), one()), (reg_t(8), m1()), (CONST_IDX, one())]);
b.lin(ROW_PAT + 1, A1, B1, sel(1), &[(reg_t(7), one()), (reg_t(4), m1())]);
b.lin(ROW_PAT + 2, A1, B1, sel(4), &[(reg_t(10), one())]);
b.put(F, ROW_PAT + 2, &[(sel(4), one())]);
b.put(G, ROW_PAT + 2, &[(reg_t(10), m1())]);
b.put(H, ROW_PAT + 2, &[(reg_t(4), one())]);
b.put(I, ROW_PAT + 2, &[(reg_t(5), one())]);
b.prod(ROW_PAT + 3, CDE1, sel(4), &[(reg_t(4), one())], &[(CONST_IDX, one()), (reg_t(10), m1())]);
b.lin(ROW_PAT + 3, A1, B1, IDX_KAPPA, &[(reg_t1(14), one()), (reg_t(14), m1()), (CONST_IDX, m1())]);
b.lin(ROW_PAT + 4, A1, B1, sel(5), &[(reg_t(6), one()), (reg_t(4), m1()), (reg_t(5), m1())]);
b.lin(ROW_PAT + 5, A1, B1, sel(6), &[(reg_t(6), one()), (reg_t(4), m1()), (reg_t(5), one())]);
b.lin(ROW_PAT + 6, A1, B1, sel(7), &[(reg_t(6), one())]);
b.prod(ROW_PAT + 6, CDE1, sel(7), &[(reg_t(4), m1())], &[(reg_t(5), one())]);
b.put(F, ROW_PAT + 7, &[(sel(8), one())]);
b.put(G, ROW_PAT + 7, &[(reg_t(6), one())]);
b.put(H, ROW_PAT + 7, &[(reg_t(6), one())]);
b.put(I, ROW_PAT + 7, &[(reg_t(4), one())]);
b.prod(ROW_PAT + 7, CDE1, sel(8), &[(reg_t(6), one())], &[(CONST_IDX, m1())]);
let diff = [(reg_t(4), one()), (reg_t(5), m1())];
let one_m_r6 = [(CONST_IDX, one()), (reg_t(6), m1())];
b.prod(ROW_PAT + 8, CDE1, sel(9), &diff, &one_m_r6);
b.prod(ROW_PAT + 9, CDE1, sel(9), &[(reg_t(6), one())], &one_m_r6);
b.prod(ROW_PAT + 10, CDE1, sel(9), &diff, &[(reg_t(7), one())]);
b.lin(ROW_PAT + 10, A1, B1, sel(9), &[(reg_t(6), m1())]);
b.lin(ROW_PAT + 8, A1, B1, sel(16), &[(reg_t(6), one())]);
let r10 = [(reg_t(10), one())];
let r11 = [(reg_t(11), one())];
let r10_m1 = [(reg_t(10), one()), (CONST_IDX, m1())];
let r11_m1 = [(reg_t(11), one()), (CONST_IDX, m1())];
b.prod(ROW_PAT + 11, CDE1, sel(10), &r10, &r10_m1);
b.prod(ROW_PAT + 14, CDE1, sel(10), &r11, &r11_m1);
b.lin(ROW_PAT + 12, A1, B1, sel(11), &[(reg_t(10), one()), (reg_t(11), one()), (reg_t(12), m1())]);
b.prod(ROW_PAT + 12, CDE1, sel(11), &[(reg_t(10), neg(Goldilocks::new(2)))], &r11);
b.prod(ROW_PAT + 1, CDE1, sel(11), &r10, &r10_m1);
b.prod(ROW_PAT + 4, CDE1, sel(11), &r11, &r11_m1);
b.prod(ROW_PAT + 13, CDE1, sel(12), &r10, &r11);
b.lin(ROW_PAT + 13, A1, B1, sel(12), &[(reg_t(12), m1())]);
b.prod(ROW_PAT + 5, CDE1, sel(12), &r10, &r10_m1);
b.prod(ROW_PAT + 15, CDE1, sel(12), &r11, &r11_m1);
b.lin(ROW_PAT + 14, A1, B1, sel(13), &[(reg_t(10), one()), (reg_t(12), one()), (CONST_IDX, m1())]);
b.lin(ROW_PAT + 7, A1, B1, sel(13), &r11);
b.lin(ROW_PAT + 15, A1, B1, sel(14), &[(reg_t(12), one()), (reg_t(11), m1())]);
CCSInstance {
matrices: b.m,
multisets: vec![
vec![A0, B0],
vec![C0, D0, E0],
vec![A1, B1],
vec![C1, D1, E1],
vec![F, G, H, I],
vec![P],
],
coeffs: vec![Goldilocks::ONE; 6],
num_rows: NUM_ROWS,
num_cols: Z_LEN,
}
}
pub fn universal_ccs() -> &'static CCSInstance {
static CELL: OnceLock<CCSInstance> = OnceLock::new();
CELL.get_or_init(build)
}
#[derive(Clone, Copy, Debug, Default)]
pub struct Capacity {
pub k: [Goldilocks; 8],
pub k1: [Goldilocks; 8],
}
pub fn universal_witness(
regs_t: &[Goldilocks; 16],
regs_t1: &[Goldilocks; 16],
caps: Option<&Capacity>,
) -> CCSWitness {
let mut z = vec![Goldilocks::ZERO; Z_LEN];
z[..16].copy_from_slice(regs_t);
z[16..32].copy_from_slice(regs_t1);
z[CONST_IDX] = Goldilocks::ONE;
let tag = regs_t[0].canonicalize().as_u64();
if (tag as usize) < NUM_PATTERNS {
z[sel(tag as usize)] = Goldilocks::ONE;
}
if tag == 15 {
let k = regs_t[14].canonicalize().as_u64() as usize;
if k < NUM_ROUNDS {
z[round_sel(k)] = Goldilocks::ONE;
}
if k < NUM_ROUNDS - 1 {
z[IDX_KAPPA] = Goldilocks::ONE;
}
if let Some(c) = caps {
z[cap_k(0)..cap_k(8)].copy_from_slice(&c.k);
z[cap_k1(0)..cap_k1(8)].copy_from_slice(&c.k1);
}
if (PARTIAL_FIRST..PARTIAL_FIRST + NUM_PARTIAL).contains(&k) {
let rc = partial_rc(k - PARTIAL_FIRST);
z[IDX_PI] = Goldilocks::ONE;
z[IDX_RC] = rc;
z[IDX_Y] = (regs_t[4] + rc).inv();
}
}
CCSWitness { z }
}
#[cfg(test)]
pub(crate) fn test_witness(vals: &[(usize, u64)]) -> CCSWitness {
let mut regs = [[Goldilocks::ZERO; 16]; 2];
for &(idx, v) in vals {
regs[idx / 16][idx % 16] = Goldilocks::new(v);
}
universal_witness(®s[0], ®s[1], None)
}
#[cfg(test)]
mod tests;
pub fn poseidon_regs(state: &[Goldilocks; 16], k: usize) -> [Goldilocks; 16] {
let mut r = [Goldilocks::ZERO; 16];
r[0] = Goldilocks::new(15);
r[4..8].copy_from_slice(&state[0..4]);
r[10..14].copy_from_slice(&state[4..8]);
r[14] = Goldilocks::new(k as u64);
r
}
pub fn capacity(state_k: &[Goldilocks; 16], state_k1: &[Goldilocks; 16]) -> Capacity {
let mut c = Capacity::default();
c.k.copy_from_slice(&state_k[8..16]);
c.k1.copy_from_slice(&state_k1[8..16]);
c
}
pub fn poseidon_witness(state_k: &[Goldilocks; 16], state_k1: &[Goldilocks; 16], k: usize) -> CCSWitness {
let caps = capacity(state_k, state_k1);
universal_witness(&poseidon_regs(state_k, k), &poseidon_regs(state_k1, k + 1), Some(&caps))
}