use crate::derive_support::PolydatSetup;
pub const MULT: u64 = 6364136223846793005;
#[inline]
pub fn pcg_output(state: u64) -> u64 {
let word = ((state >> ((state >> 59) + 5)) ^ state).wrapping_mul(12605985483714917081);
(word >> 43) ^ word
}
#[inline]
pub 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)
}
pub 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 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,
}
}
#[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 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;
}
}
}