use crate::ast::CompiledU64Op;
use crate::derive_support::Const;
#[crate::polydat_node(category = Arithmetic)]
fn add(input: u64, addend: Const<u64>) -> u64 {
input.wrapping_add(*addend)
}
#[crate::polydat_node(category = Arithmetic)]
fn mul(input: u64, factor: Const<u64>) -> u64 {
input.wrapping_mul(*factor)
}
#[crate::polydat_node(category = Arithmetic)]
fn div(input: u64, divisor: Const<u64>) -> u64 {
input / *divisor
}
#[crate::polydat_node(category = Arithmetic)]
fn r#mod(input: u64, modulus: Const<u64>) -> u64 {
input % *modulus
}
#[crate::polydat_node(category = Arithmetic)]
fn mod_wire(input: u64, #[constraint(NonZeroU64)] divisor: u64) -> u64 {
input % divisor
}
#[crate::polydat_node(category = Arithmetic)]
fn div_wire(input: u64, #[constraint(NonZeroU64)] divisor: u64) -> u64 {
input / divisor
}
#[crate::polydat_node(category = Arithmetic)]
fn ceil_to_multiple(value: u64, multiple: u64) -> u64 {
if multiple == 0 {
value
} else {
value.div_ceil(multiple).saturating_mul(multiple)
}
}
#[crate::polydat_node(category = Arithmetic)]
fn multiples_at_least(value: u64, multiple: u64) -> u64 {
if multiple == 0 { 0 } else { value.div_ceil(multiple) }
}
#[crate::polydat_node(category = Arithmetic)]
fn set_or_get(current: u64, fallback: u64) -> u64 {
if current == 0 { fallback } else { current }
}
#[crate::polydat_node(category = Arithmetic)]
fn clamp(input: u64, min: Const<u64>, max: Const<u64>) -> u64 {
input.clamp(*min, *max)
}
fn mixed_radix_jit(node: &MixedRadix) -> CompiledU64Op {
let radixes = node.radixes.clone();
Box::new(move |inputs, outputs| {
let mut remainder = inputs[0];
for (i, &radix) in radixes.iter().enumerate() {
if radix == 0 {
outputs[i] = remainder;
remainder = 0;
} else {
outputs[i] = remainder % radix;
remainder /= radix;
}
}
})
}
fn mixed_radix_jit_constants(node: &MixedRadix) -> Vec<u64> {
node.radixes.clone()
}
#[crate::polydat_node(
category = Arithmetic,
compiled_u64 = mixed_radix_jit,
jit_constants = mixed_radix_jit_constants,
)]
fn mixed_radix(
input: u64,
radixes: crate::derive_support::Const<Vec<u64>>,
) -> crate::derive_support::DynamicOutputs<u64> {
let mut remainder = input;
let mut result = Vec::with_capacity(radixes.len());
for &radix in radixes.iter() {
if radix == 0 {
result.push(remainder);
remainder = 0;
} else {
result.push(remainder % radix);
remainder /= radix;
}
}
crate::derive_support::DynamicOutputs(result)
}
#[crate::polydat_node(category = Variadic, identity = 0u64, commutativity = AllCommutative)]
fn sum(values: &[u64]) -> u64 {
values.iter().fold(0u64, |a, b| a.wrapping_add(*b))
}
#[crate::polydat_node(category = Variadic, identity = 1u64, commutativity = AllCommutative)]
fn product(values: &[u64]) -> u64 {
values.iter().fold(1u64, |a, b| a.wrapping_mul(*b))
}
#[crate::polydat_node(category = Variadic, identity = u64::MAX, commutativity = AllCommutative)]
fn min(values: &[u64]) -> u64 {
values.iter().copied().fold(u64::MAX, std::cmp::min)
}
#[crate::polydat_node(category = Variadic, identity = 0u64, commutativity = AllCommutative)]
fn max(values: &[u64]) -> u64 {
values.iter().copied().fold(0u64, std::cmp::max)
}
#[crate::polydat_node(category = Arithmetic)]
fn interleave(a: u64, b: u64) -> u64 {
let mut result: u64 = 0;
for i in 0..32 {
result |= ((a >> i) & 1) << (2 * i);
result |= ((b >> i) & 1) << (2 * i + 1);
}
result
}
use crate::dsl::registry::FuncSig;
pub fn signatures() -> &'static [FuncSig] {
&[]
}
pub(crate) fn build_node(name: &str, _wires: &[crate::compile::assembly::WireRef], _wire_types: &[crate::ast::PortType], consts: &[crate::dsl::factory::ConstArg]) -> Option<Result<Box<dyn crate::ast::PolydatNode>, String>> {
let _ = (name, consts);
None
}
pub(crate) fn validate_node(
name: &str,
consts: &[crate::dsl::factory::ConstArg],
) -> Result<(), String> {
match name {
"mixed_radix" => {
for (i, c) in consts.iter().enumerate().take(consts.len().saturating_sub(1)) {
if c.as_u64() == 0 {
return Err(format!("radix {i} must be non-zero"));
}
}
Ok(())
}
_ => Ok(()),
}
}
crate::register_nodes!(signatures, build_node, validate_node);
#[cfg(test)]
mod tests {
use super::*;
use crate::ast::{PolydatNode, Value};
#[test]
fn add_wrapping() {
let node = Add::new(10);
let mut out = [Value::None];
node.eval(&[Value::U64(5)], &mut out);
assert_eq!(out[0].as_u64(), 15);
}
#[test]
fn mod_basic() {
let node = Mod::new(100);
let mut out = [Value::None];
node.eval(&[Value::U64(542)], &mut out);
assert_eq!(out[0].as_u64(), 42);
}
#[test]
fn mixed_radix_decompose() {
let node = MixedRadix::new(vec![100, 1000, 0]);
let mut out = [Value::None, Value::None, Value::None];
node.eval(&[Value::U64(4_201_337)], &mut out);
assert_eq!(out[0].as_u64(), 37);
assert_eq!(out[1].as_u64(), 13);
assert_eq!(out[2].as_u64(), 42);
}
#[test]
fn mixed_radix_cartesian() {
let node = MixedRadix::new(vec![100, 1000, 0]);
let mut out = [Value::None, Value::None, Value::None];
node.eval(&[Value::U64(0)], &mut out);
assert_eq!(out[0].as_u64(), 0);
assert_eq!(out[1].as_u64(), 0);
assert_eq!(out[2].as_u64(), 0);
node.eval(&[Value::U64(100_000)], &mut out);
assert_eq!(out[0].as_u64(), 0);
assert_eq!(out[1].as_u64(), 0);
assert_eq!(out[2].as_u64(), 1);
}
#[test]
fn interleave_basic() {
let node = Interleave::new();
let mut out = [Value::None];
node.eval(&[Value::U64(0b101), Value::U64(0b010)], &mut out);
assert_eq!(out[0].as_u64(), 0b01_10_01);
}
#[test]
fn div_basic() {
let node = Div::new(100);
let mut out = [Value::None];
node.eval(&[Value::U64(4_201_337)], &mut out);
assert_eq!(out[0].as_u64(), 42013);
}
#[test]
fn sum_variadic() {
let node = Sum::new(0);
let mut out = [Value::None];
node.eval(&[], &mut out);
assert_eq!(out[0].as_u64(), 0);
let node = Sum::new(1);
node.eval(&[Value::U64(42)], &mut out);
assert_eq!(out[0].as_u64(), 42);
let node = Sum::new(3);
node.eval(&[Value::U64(10), Value::U64(20), Value::U64(30)], &mut out);
assert_eq!(out[0].as_u64(), 60);
}
#[test]
fn product_variadic() {
let node = Product::new(0);
let mut out = [Value::None];
node.eval(&[], &mut out);
assert_eq!(out[0].as_u64(), 1);
let node = Product::new(1);
node.eval(&[Value::U64(7)], &mut out);
assert_eq!(out[0].as_u64(), 7);
let node = Product::new(3);
node.eval(&[Value::U64(2), Value::U64(3), Value::U64(7)], &mut out);
assert_eq!(out[0].as_u64(), 42);
}
#[test]
fn min_variadic() {
let node = Min::new(0);
let mut out = [Value::None];
node.eval(&[], &mut out);
assert_eq!(out[0].as_u64(), u64::MAX);
let node = Min::new(3);
node.eval(&[Value::U64(50), Value::U64(10), Value::U64(30)], &mut out);
assert_eq!(out[0].as_u64(), 10);
}
#[test]
fn max_variadic() {
let node = Max::new(0);
let mut out = [Value::None];
node.eval(&[], &mut out);
assert_eq!(out[0].as_u64(), 0);
let node = Max::new(3);
node.eval(&[Value::U64(50), Value::U64(10), Value::U64(30)], &mut out);
assert_eq!(out[0].as_u64(), 50);
}
#[test]
fn slot_constants_match_jit_constants() {
use crate::ast::PolydatNode;
let nodes: Vec<Box<dyn PolydatNode>> = vec![
Box::new(Add::new(42)),
Box::new(Mul::new(7)),
Box::new(Div::new(100)),
Box::new(Mod::new(256)),
Box::new(Clamp::new(10, 90)),
Box::new(MixedRadix::new(vec![100, 1000, 0])),
];
for node in &nodes {
let from_trait = node.jit_constants();
let from_slots = node.meta().jit_constants_from_slots();
assert_eq!(
from_trait, from_slots,
"constant mismatch for node '{}': trait={from_trait:?}, slots={from_slots:?}",
node.meta().name,
);
}
}
fn run_binary(node: &dyn PolydatNode, a: u64, b: u64) -> u64 {
let mut out = [Value::None];
node.eval(&[Value::U64(a), Value::U64(b)], &mut out);
out[0].as_u64()
}
#[test]
fn ceil_to_multiple_returns_value_when_already_a_multiple() {
let n = CeilToMultiple::default();
assert_eq!(run_binary(&n, 800, 100), 800);
}
#[test]
fn ceil_to_multiple_rounds_up_to_next_boundary() {
let n = CeilToMultiple::default();
assert_eq!(run_binary(&n, 801, 100), 900);
}
#[test]
fn ceil_to_multiple_zero_value_is_zero() {
let n = CeilToMultiple::default();
assert_eq!(run_binary(&n, 0, 100), 0);
}
#[test]
fn ceil_to_multiple_below_one_multiple_rounds_to_multiple() {
let n = CeilToMultiple::default();
assert_eq!(run_binary(&n, 50, 100), 100);
assert_eq!(run_binary(&n, 1, 100), 100);
}
#[test]
fn ceil_to_multiple_zero_multiple_is_soft_no_op() {
let n = CeilToMultiple::default();
assert_eq!(run_binary(&n, 42, 0), 42,
"multiple=0 must not trap; passes value through");
}
#[test]
fn multiples_at_least_exact_division() {
let n = MultiplesAtLeast::default();
assert_eq!(run_binary(&n, 800, 100), 8);
}
#[test]
fn multiples_at_least_rounds_up_partial() {
let n = MultiplesAtLeast::default();
assert_eq!(run_binary(&n, 801, 100), 9);
assert_eq!(run_binary(&n, 1, 100), 1);
}
#[test]
fn multiples_at_least_zero_value_is_zero() {
let n = MultiplesAtLeast::default();
assert_eq!(run_binary(&n, 0, 100), 0);
}
#[test]
fn multiples_at_least_zero_multiple_is_zero() {
let n = MultiplesAtLeast::default();
assert_eq!(run_binary(&n, 42, 0), 0);
}
#[test]
fn set_or_get_returns_current_when_non_zero() {
let n = SetOrGet::default();
assert_eq!(run_binary(&n, 7, 99), 7);
assert_eq!(run_binary(&n, u64::MAX, 99), u64::MAX);
}
#[test]
fn set_or_get_returns_fallback_when_current_is_zero() {
let n = SetOrGet::default();
assert_eq!(run_binary(&n, 0, 99), 99);
}
#[test]
fn set_or_get_zero_fallback_is_zero() {
let n = SetOrGet::default();
assert_eq!(run_binary(&n, 0, 0), 0);
}
#[test]
fn set_or_get_idempotent_on_already_set() {
let n = SetOrGet::default();
for v in [1u64, 42, 1000, u64::MAX] {
assert_eq!(run_binary(&n, v, 999), v);
}
}
#[test]
fn ceil_to_multiple_and_count_satisfy_invariant() {
let ceil = CeilToMultiple::default();
let count = MultiplesAtLeast::default();
for (v, m) in [(0u64, 100), (1, 100), (50, 100), (100, 100),
(101, 100), (10000, 7), (10000, 64), (12345, 256)] {
let c_val = run_binary(&ceil, v, m);
let n_val = run_binary(&count, v, m);
assert_eq!(c_val, n_val * m,
"invariant violated for (v={v}, m={m}): ceil={c_val}, count={n_val}");
}
}
#[test]
fn slot_wire_inputs_match_inputs() {
use crate::ast::PolydatNode;
let nodes: Vec<Box<dyn PolydatNode>> = vec![
Box::new(Add::new(0)),
Box::new(Mod::new(1)),
Box::new(Sum::new(3)),
Box::new(Product::new(2)),
Box::new(Interleave::new()),
Box::new(MixedRadix::new(vec![10, 20])),
Box::new(CeilToMultiple::default()),
Box::new(MultiplesAtLeast::default()),
Box::new(SetOrGet::default()),
];
for node in &nodes {
let old_count = node.meta().wire_inputs().len();
let new_count = node.meta().wire_inputs().len();
assert_eq!(
old_count, new_count,
"wire input count mismatch for '{}': inputs={old_count}, wire_inputs()={new_count}",
node.meta().name,
);
}
}
}