use crate::ast::CompiledU64Op;
#[cfg(test)]
use crate::ast::{PolydatNode, Value};
use crate::derive_support::{Const, PolydatSetup};
const MULT: u64 = 6364136223846793005;
#[inline]
pub(crate) fn pcg_output(state: u64) -> u64 {
let word = ((state >> ((state >> 59) + 5)) ^ state)
.wrapping_mul(12605985483714917081);
(word >> 43) ^ word
}
#[inline]
pub(crate) fn pcg_seek(seed: u64, inc: u64, position: u64) -> u64 {
let mut cur_mult = MULT;
let mut cur_plus = inc;
let mut acc_mult: u64 = 1;
let mut acc_plus: u64 = 0;
let mut delta = position;
while delta > 0 {
if delta & 1 != 0 {
acc_mult = acc_mult.wrapping_mul(cur_mult);
acc_plus = acc_plus.wrapping_mul(cur_mult).wrapping_add(cur_plus);
}
cur_plus = cur_mult.wrapping_add(1).wrapping_mul(cur_plus);
cur_mult = cur_mult.wrapping_mul(cur_mult);
delta >>= 1;
}
let state = acc_mult.wrapping_mul(seed).wrapping_add(acc_plus);
pcg_output(state)
}
#[crate::polydat_node(category = Permutation)]
fn pcg(
input: u64,
#[poly_default(0u64)] seed: Const<u64>,
#[poly_default(0u64)] stream: Const<u64>,
) -> u64 {
let inc = 2u64.wrapping_mul(*stream).wrapping_add(1);
pcg_seek(*seed, inc, input)
}
#[crate::polydat_node(category = Permutation)]
fn pcg_stream(input: u64, stream: u64, #[poly_default(0u64)] seed: Const<u64>) -> u64 {
let inc = 2u64.wrapping_mul(stream).wrapping_add(1);
pcg_seek(*seed, inc, input)
}
const FEISTEL_ROUNDS: usize = 6;
pub struct CycleWalkState {
pub half_bits: u32,
pub half_mask: u64,
pub inc: u64,
pub round_keys: [u64; FEISTEL_ROUNDS],
}
impl PolydatSetup for CycleWalkState {}
pub(crate) fn build_cycle_walk_state(range: u64, seed: u64, stream: u64) -> CycleWalkState {
assert!(range > 0, "CycleWalk range must be > 0");
let inc = 2u64.wrapping_mul(stream).wrapping_add(1);
let min_bits = if range <= 1 {
2 } else {
let b = 64 - (range - 1).leading_zeros();
if !b.is_multiple_of(2) { b + 1 } else { b.max(2) }
};
let half_bits = min_bits / 2;
let half_mask = (1u64 << half_bits) - 1;
let mut round_keys = [0u64; FEISTEL_ROUNDS];
for (i, key) in round_keys.iter_mut().enumerate() {
*key = pcg_seek(seed, inc, i as u64 + 1_000_000_000);
}
CycleWalkState { half_bits, half_mask, inc, round_keys }
}
fn cycle_walk_jit(node: &CycleWalk) -> CompiledU64Op {
let range = node.range;
let half_bits = node.state.half_bits;
let half_mask = node.state.half_mask;
let round_keys = node.state.round_keys;
Box::new(move |inputs, outputs| {
outputs[0] = cycle_walk_inner(inputs[0], range, half_bits, half_mask, &round_keys);
})
}
fn cycle_walk_jit_constants(node: &CycleWalk) -> Vec<u64> {
vec![node.range, node.seed, node.state.inc]
}
#[crate::polydat_node(
category = Permutation,
compiled_u64 = cycle_walk_jit,
jit_constants = cycle_walk_jit_constants,
)]
fn cycle_walk(
position: u64,
range: Const<u64>,
#[poly_default(0u64)] seed: Const<u64>,
#[poly_default(0u64)] stream: Const<u64>,
#[poly_const(build_cycle_walk_state, from = (range, seed, stream))]
state: &CycleWalkState,
) -> u64 {
let _ = seed;
let _ = stream;
cycle_walk_inner(
position,
*range,
state.half_bits,
state.half_mask,
&state.round_keys,
)
}
#[inline]
fn feistel_round_fn(half: u64, round_key: u64) -> u64 {
let x = half.wrapping_mul(0x9E3779B97F4A7C15).wrapping_add(round_key);
let x = ((x >> 32) ^ x).wrapping_mul(0xD6E8FEB86659FD93);
(x >> 32) ^ x
}
#[inline]
fn feistel_encrypt(
value: u64,
half_bits: u32,
half_mask: u64,
round_keys: &[u64; FEISTEL_ROUNDS],
) -> u64 {
let mut left = (value >> half_bits) & half_mask;
let mut right = value & half_mask;
for key in round_keys.iter() {
let new_right = left ^ (feistel_round_fn(right, *key) & half_mask);
left = right;
right = new_right;
}
(left << half_bits) | right
}
#[inline]
pub(crate) fn cycle_walk_inner(
mut value: u64,
range: u64,
half_bits: u32,
half_mask: u64,
round_keys: &[u64; FEISTEL_ROUNDS],
) -> u64 {
if range == 1 {
return 0;
}
value %= range;
loop {
value = feistel_encrypt(value, half_bits, half_mask, round_keys);
if value < range {
return value;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashSet;
#[test]
fn pcg_output_deterministic() {
let a = pcg_output(123456789);
let b = pcg_output(123456789);
assert_eq!(a, b);
}
#[test]
fn pcg_seek_position_zero_vs_one() {
let seed = 42u64;
let inc = 1u64; let v0 = pcg_seek(seed, inc, 0);
let v1 = pcg_seek(seed, inc, 1);
assert_ne!(v0, v1, "different positions must produce different values");
}
#[test]
fn pcg_seek_deterministic() {
let seed = 0xDEAD_BEEF;
let inc = 3;
let a = pcg_seek(seed, inc, 1000);
let b = pcg_seek(seed, inc, 1000);
assert_eq!(a, b);
}
#[test]
fn pcg_seek_sequential_matches_step() {
let seed = 77u64;
let inc = 5u64;
let n = 50u64;
let mut state = seed;
for _ in 0..n {
state = state.wrapping_mul(MULT).wrapping_add(inc);
}
let stepped = pcg_output(state);
let seeked = pcg_seek(seed, inc, n);
assert_eq!(stepped, seeked,
"seek({n}) must match {n} sequential LCG steps");
}
#[test]
fn pcg_node_deterministic() {
let node = Pcg::new(42, 0);
let mut out = [Value::None];
node.eval(&[Value::U64(100)], &mut out);
let first = out[0].as_u64();
node.eval(&[Value::U64(100)], &mut out);
assert_eq!(first, out[0].as_u64(), "same position must give same result");
}
#[test]
fn pcg_node_different_positions() {
let node = Pcg::new(42, 0);
let mut out1 = [Value::None];
let mut out2 = [Value::None];
node.eval(&[Value::U64(0)], &mut out1);
node.eval(&[Value::U64(1)], &mut out2);
assert_ne!(out1[0].as_u64(), out2[0].as_u64());
}
#[test]
fn pcg_node_different_seeds() {
let a = Pcg::new(1, 0);
let b = Pcg::new(2, 0);
let mut out_a = [Value::None];
let mut out_b = [Value::None];
a.eval(&[Value::U64(50)], &mut out_a);
b.eval(&[Value::U64(50)], &mut out_b);
assert_ne!(out_a[0].as_u64(), out_b[0].as_u64(),
"different seeds should produce different values");
}
#[test]
fn pcg_node_different_streams() {
let a = Pcg::new(42, 0);
let b = Pcg::new(42, 1);
let mut out_a = [Value::None];
let mut out_b = [Value::None];
a.eval(&[Value::U64(50)], &mut out_a);
b.eval(&[Value::U64(50)], &mut out_b);
assert_ne!(out_a[0].as_u64(), out_b[0].as_u64(),
"different streams should produce different values");
}
#[test]
fn pcg_compiled_matches_eval() {
let node = Pcg::new(99, 7);
let compiled = node.compiled_u64().expect("Pcg must provide compiled_u64");
for pos in 0..100u64 {
let mut eval_out = [Value::None];
node.eval(&[Value::U64(pos)], &mut eval_out);
let mut comp_out = [0u64];
compiled(&[pos], &mut comp_out);
assert_eq!(eval_out[0].as_u64(), comp_out[0],
"compiled and eval must agree at position {pos}");
}
}
#[test]
fn pcg_jit_constants() {
let node = Pcg::new(42, 7);
let consts = node.jit_constants();
assert_eq!(consts.len(), 2);
assert_eq!(consts[0], 42, "first constant is seed");
assert_eq!(consts[1], 7, "second constant is stream (inc = 2*stream+1 derived in body)");
}
#[test]
fn pcg_stream_deterministic() {
let node = PcgStream::new(42);
let mut out = [Value::None];
node.eval(&[Value::U64(100), Value::U64(3)], &mut out);
let first = out[0].as_u64();
node.eval(&[Value::U64(100), Value::U64(3)], &mut out);
assert_eq!(first, out[0].as_u64());
}
#[test]
fn pcg_stream_independence() {
let node = PcgStream::new(42);
let mut out_a = [Value::None];
let mut out_b = [Value::None];
node.eval(&[Value::U64(50), Value::U64(0)], &mut out_a);
node.eval(&[Value::U64(50), Value::U64(1)], &mut out_b);
assert_ne!(out_a[0].as_u64(), out_b[0].as_u64(),
"different stream_ids should produce different values");
}
#[test]
fn pcg_stream_matches_fixed_pcg() {
let fixed = Pcg::new(42, 5);
let dynamic = PcgStream::new(42);
for pos in 0..50u64 {
let mut f_out = [Value::None];
let mut d_out = [Value::None];
fixed.eval(&[Value::U64(pos)], &mut f_out);
dynamic.eval(&[Value::U64(pos), Value::U64(5)], &mut d_out);
assert_eq!(f_out[0].as_u64(), d_out[0].as_u64(),
"PcgStream must match Pcg for same seed/stream at position {pos}");
}
}
#[test]
fn pcg_stream_compiled_matches_eval() {
let node = PcgStream::new(99);
let compiled = node.compiled_u64().expect("PcgStream must provide compiled_u64");
for pos in 0..50u64 {
for stream in 0..5u64 {
let mut eval_out = [Value::None];
node.eval(&[Value::U64(pos), Value::U64(stream)], &mut eval_out);
let mut comp_out = [0u64];
compiled(&[pos, stream], &mut comp_out);
assert_eq!(eval_out[0].as_u64(), comp_out[0],
"compiled and eval must agree at pos={pos}, stream={stream}");
}
}
}
#[test]
fn cycle_walk_bounded() {
let node = CycleWalk::new(100, 42, 0);
let mut out = [Value::None];
for i in 0..200u64 {
node.eval(&[Value::U64(i)], &mut out);
assert!(out[0].as_u64() < 100, "output {} >= range 100", out[0].as_u64());
}
}
#[test]
fn cycle_walk_deterministic() {
let node = CycleWalk::new(1000, 42, 0);
let mut out = [Value::None];
node.eval(&[Value::U64(77)], &mut out);
let first = out[0].as_u64();
node.eval(&[Value::U64(77)], &mut out);
assert_eq!(first, out[0].as_u64());
}
#[test]
fn cycle_walk_bijective_small() {
let range = 50u64;
let node = CycleWalk::new(range, 42, 0);
let mut seen = HashSet::new();
let mut out = [Value::None];
for i in 0..range {
node.eval(&[Value::U64(i)], &mut out);
let v = out[0].as_u64();
assert!(v < range, "output {v} out of range [0, {range})");
assert!(seen.insert(v), "duplicate output {v} at position {i}");
}
assert_eq!(seen.len(), range as usize,
"must produce exactly {range} distinct values");
}
#[test]
fn cycle_walk_bijective_power_of_two() {
let range = 64u64;
let node = CycleWalk::new(range, 123, 7);
let mut seen = HashSet::new();
let mut out = [Value::None];
for i in 0..range {
node.eval(&[Value::U64(i)], &mut out);
let v = out[0].as_u64();
assert!(v < range);
assert!(seen.insert(v), "duplicate at {i}");
}
assert_eq!(seen.len(), range as usize);
}
#[test]
fn cycle_walk_compiled_matches_eval() {
let node = CycleWalk::new(200, 42, 3);
let compiled = node.compiled_u64().expect("CycleWalk must provide compiled_u64");
for pos in 0..200u64 {
let mut eval_out = [Value::None];
node.eval(&[Value::U64(pos)], &mut eval_out);
let mut comp_out = [0u64];
compiled(&[pos], &mut comp_out);
assert_eq!(eval_out[0].as_u64(), comp_out[0],
"compiled and eval must agree at position {pos}");
}
}
#[test]
fn cycle_walk_jit_constants() {
let node = CycleWalk::new(500, 42, 7);
let consts = node.jit_constants();
assert_eq!(consts.len(), 3);
assert_eq!(consts[0], 500, "first constant is range");
assert_eq!(consts[1], 42, "second constant is seed");
assert_eq!(consts[2], 2 * 7 + 1, "third constant is inc");
}
#[test]
#[should_panic(expected = "range must be > 0")]
fn cycle_walk_zero_range_panics() {
CycleWalk::new(0, 42, 0);
}
#[test]
fn cycle_walk_range_one() {
let node = CycleWalk::new(1, 42, 0);
let mut out = [Value::None];
for i in 0..10u64 {
node.eval(&[Value::U64(i)], &mut out);
assert_eq!(out[0].as_u64(), 0);
}
}
}