Skip to main content

vyre_primitives/math/
bigint_add_carry.rs

1//! `bigint_add_carry`  -  multi-limb big-integer addition with explicit
2//! carry-out propagation, packed as one u32-limb per element.
3//!
4//! Op id: `vyre-primitives::math::bigint_add_carry`. Soundness: `Exact` over
5//! `(a + b) mod 2^(32 * limb_count)` with the high carry-out emitted as a
6//! separate scalar. The CPU reference at the bottom of this file is the
7//! contract; the GPU `Program` matches it lane-for-lane.
8//!
9//! ## Why it matters
10//!
11//! Public-key crypto (RSA, ECDSA, post-quantum lattices), digital-signature
12//! verification, and arbitrary-precision integer math all bottom out into
13//! ripple-carry addition over 256-bit / 512-bit / 4096-bit operands. Doing
14//! this on GPU naively serializes all carries through a single thread: each
15//! limb depends on the carry from the limb below. This primitive ships the
16//! foundation: a load-and-add wave that emits per-limb sums + per-limb
17//! carry-out booleans. A carry-fix wave sweeps the per-limb carry stream
18//! into a final answer.
19//!
20//! The output layout is the canonical "split-carry" form expected by every
21//! known parallel bigint adder (Brent-Kung, Kogge-Stone, Sklansky). Once you
22//! have `(sum_no_carry[i], carry[i])` you can finish in O(log n) prefix-scan
23//! depth instead of O(n) ripple. This module emits the first half; the
24//! prefix-scan finish is in `prefix_scan` (#5).
25//!
26//! ## Wire layout
27//!
28//! Inputs:
29//!   - `a`  -  limb_count u32 limbs, little-endian (limb 0 = LSB).
30//!   - `b`  -  limb_count u32 limbs, little-endian.
31//!
32//! Outputs:
33//!   - `sum_partial`  -  limb_count u32 limbs: `(a[i] + b[i]) mod 2^32`.
34//!   - `carry_partial`  -  limb_count u32 limbs (each is 0 or 1): the
35//!     carry-out of `a[i] + b[i]`. Bit `i` of the final carry-resolved sum
36//!     comes from `sum_partial[i] + carry_in[i]` where `carry_in[i]` is
37//!     the prefix-or-style fold of `carry_partial[0..i]` adjusted for
38//!     "carry-generate" from the partial sum overflowing.
39//!
40//! This module is the load-and-half-add primitive used by the parallel-prefix
41//! carry resolver.
42
43use std::sync::Arc;
44
45use vyre_foundation::ir::model::expr::Ident;
46use vyre_foundation::ir::{BufferAccess, BufferDecl, DataType, Expr, Node, Program};
47
48/// Canonical op id for region-chain audits and bench attribution.
49pub const OP_ID: &str = "vyre-primitives::math::bigint_add_carry";
50
51/// Canonical binding indices.
52pub const BINDING_A_IN: u32 = 0;
53/// `b` operand binding.
54pub const BINDING_B_IN: u32 = 1;
55/// `sum_partial` output binding.
56pub const BINDING_SUM_PARTIAL_OUT: u32 = 2;
57/// `carry_partial` output binding (one u32 per limb, value is 0 or 1).
58pub const BINDING_CARRY_PARTIAL_OUT: u32 = 3;
59
60/// One lane per bigint limb in the split add-carry pass.
61pub const BIGINT_ADD_CARRY_WORKGROUP_SIZE: [u32; 3] = [256, 1, 1];
62
63/// Dispatch grid that covers every bigint limb lane.
64#[must_use]
65pub const fn bigint_add_carry_dispatch_grid(limb_count: u32) -> [u32; 3] {
66    let lanes_per_block = BIGINT_ADD_CARRY_WORKGROUP_SIZE[0];
67    let full_blocks = limb_count / lanes_per_block;
68    let tail_block = if limb_count % lanes_per_block == 0 {
69        0
70    } else {
71        1
72    };
73    let blocks = full_blocks + tail_block;
74    [if blocks == 0 { 1 } else { blocks }, 1, 1]
75}
76
77/// Bigint CPU-reference error.
78#[derive(Debug, Clone, PartialEq, Eq)]
79#[non_exhaustive]
80pub enum BigIntAddCarryError {
81    /// Input operands had different limb counts.
82    LimbCountMismatch {
83        /// `a` operand length.
84        a_len: usize,
85        /// `b` operand length.
86        b_len: usize,
87    },
88    /// Split carry arrays had different limb counts.
89    SplitCarryLengthMismatch {
90        /// `sum_partial` length.
91        sum_len: usize,
92        /// `carry_partial` length.
93        carry_len: usize,
94    },
95    /// Caller-owned storage could not be reserved.
96    AllocationFailed {
97        /// Operation that was reserving storage.
98        operation: &'static str,
99        /// Allocator or capacity diagnostic.
100        message: String,
101    },
102}
103
104/// Build the IR `Program` that emits `(sum_partial, carry_partial)` for
105/// a multi-limb big-integer addition.
106///
107/// One thread per limb. Each thread:
108///   1. Loads `a[gid]` and `b[gid]`.
109///   2. Computes `sum = a + b` mod 2^32 and `carry = if sum < a { 1 } else { 0 }`
110///      (canonical "carry from unsigned overflow" check).
111///   3. Stores both into the output buffers at index `gid`.
112///
113/// `limb_count` must be > 0; the workgroup size is fixed at 256 lanes.
114#[must_use]
115pub fn bigint_add_carry(limb_count: u32) -> Program {
116    if limb_count == 0 {
117        return crate::invalid_output_program(
118            OP_ID,
119            "sum_partial",
120            DataType::U32,
121            "Fix: bigint_add_carry requires limb_count > 0, got 0.".to_string(),
122        );
123    }
124
125    let body = vec![
126        Node::let_bind("limb_idx", Expr::InvocationId { axis: 0 }),
127        Node::if_then(
128            Expr::lt(Expr::var("limb_idx"), Expr::u32(limb_count)),
129            vec![
130                Node::let_bind("a_limb", Expr::load("a", Expr::var("limb_idx"))),
131                Node::let_bind("b_limb", Expr::load("b", Expr::var("limb_idx"))),
132                // sum (mod 2^32). Hardware u32 add already wraps.
133                Node::let_bind("sum", Expr::add(Expr::var("a_limb"), Expr::var("b_limb"))),
134                // carry = (sum < a_limb) ? 1 : 0    -  the canonical
135                // detect-unsigned-overflow check.
136                Node::let_bind(
137                    "carry_bool",
138                    Expr::lt(Expr::var("sum"), Expr::var("a_limb")),
139                ),
140                Node::let_bind(
141                    "carry",
142                    Expr::select(Expr::var("carry_bool"), Expr::u32(1), Expr::u32(0)),
143                ),
144                Node::store("sum_partial", Expr::var("limb_idx"), Expr::var("sum")),
145                Node::store("carry_partial", Expr::var("limb_idx"), Expr::var("carry")),
146            ],
147        ),
148    ];
149
150    let buffers = vec![
151        BufferDecl::storage("a", BINDING_A_IN, BufferAccess::ReadOnly, DataType::U32)
152            .with_count(limb_count),
153        BufferDecl::storage("b", BINDING_B_IN, BufferAccess::ReadOnly, DataType::U32)
154            .with_count(limb_count),
155        BufferDecl::storage(
156            "sum_partial",
157            BINDING_SUM_PARTIAL_OUT,
158            BufferAccess::ReadWrite,
159            DataType::U32,
160        )
161        .with_count(limb_count),
162        BufferDecl::storage(
163            "carry_partial",
164            BINDING_CARRY_PARTIAL_OUT,
165            BufferAccess::ReadWrite,
166            DataType::U32,
167        )
168        .with_count(limb_count),
169    ];
170
171    let entry = vec![Node::Region {
172        generator: Ident::from(OP_ID),
173        source_region: None,
174        body: Arc::new(body),
175    }];
176    Program::wrapped(buffers, BIGINT_ADD_CARRY_WORKGROUP_SIZE, entry)
177}
178
179/// CPU reference. Returns `(sum_partial, carry_partial)` matching the
180/// GPU `Program` lane-for-lane.
181///
182/// Each limb is added with the per-limb carry computed in isolation
183/// (no carry chaining). The downstream prefix-scan resolves the chain.
184///
185/// # Errors
186///
187/// Returns [`BigIntAddCarryError::LimbCountMismatch`] when operands have
188/// different limb counts.
189#[cfg(any(test, feature = "cpu-parity"))]
190pub fn bigint_add_carry_cpu(
191    a: &[u32],
192    b: &[u32],
193) -> Result<(Vec<u32>, Vec<u32>), BigIntAddCarryError> {
194    let mut sum_partial = Vec::with_capacity(a.len());
195    let mut carry_partial = Vec::with_capacity(a.len());
196    bigint_add_carry_cpu_into(a, b, &mut sum_partial, &mut carry_partial)?;
197    Ok((sum_partial, carry_partial))
198}
199
200/// CPU reference into caller-owned output buffers.
201///
202/// Clears `sum_partial` and `carry_partial`, then reuses their capacity.
203///
204/// # Errors
205///
206/// Returns [`BigIntAddCarryError::LimbCountMismatch`] when operands have
207/// different limb counts.
208#[cfg(any(test, feature = "cpu-parity"))]
209pub fn bigint_add_carry_cpu_into(
210    a: &[u32],
211    b: &[u32],
212    sum_partial: &mut Vec<u32>,
213    carry_partial: &mut Vec<u32>,
214) -> Result<(), BigIntAddCarryError> {
215    if a.len() != b.len() {
216        return Err(BigIntAddCarryError::LimbCountMismatch {
217            a_len: a.len(),
218            b_len: b.len(),
219        });
220    }
221    reserve_bigint_output(sum_partial, a.len(), "sum_partial")?;
222    reserve_bigint_output(carry_partial, a.len(), "carry_partial")?;
223    sum_partial.clear();
224    carry_partial.clear();
225    for (a_limb, b_limb) in a.iter().zip(b.iter()) {
226        let (sum, overflow) = a_limb.overflowing_add(*b_limb);
227        sum_partial.push(sum);
228        carry_partial.push(u32::from(overflow));
229    }
230    Ok(())
231}
232
233/// Resolve carry chain in O(n) ripple form. Used by the CPU reference
234/// to validate that the (sum_partial, carry_partial) split form composes
235/// correctly into the final big-integer sum + final carry-out.
236///
237/// Returns `(final_sum, final_carry_out)`. `final_carry_out` is 0 or 1.
238#[cfg(any(test, feature = "cpu-parity"))]
239pub fn resolve_carry_chain_cpu(
240    sum_partial: &[u32],
241    carry_partial: &[u32],
242) -> Result<(Vec<u32>, u32), BigIntAddCarryError> {
243    let mut final_sum = Vec::with_capacity(sum_partial.len());
244    let final_carry = resolve_carry_chain_cpu_into(sum_partial, carry_partial, &mut final_sum)?;
245    Ok((final_sum, final_carry))
246}
247
248/// Resolve carry chain into caller-owned output storage.
249///
250/// Clears `final_sum`, then reuses its capacity. Returns final carry-out.
251///
252/// # Errors
253///
254/// Returns [`BigIntAddCarryError::SplitCarryLengthMismatch`] when the split
255/// carry buffers have different limb counts.
256#[cfg(any(test, feature = "cpu-parity"))]
257pub fn resolve_carry_chain_cpu_into(
258    sum_partial: &[u32],
259    carry_partial: &[u32],
260    final_sum: &mut Vec<u32>,
261) -> Result<u32, BigIntAddCarryError> {
262    if sum_partial.len() != carry_partial.len() {
263        return Err(BigIntAddCarryError::SplitCarryLengthMismatch {
264            sum_len: sum_partial.len(),
265            carry_len: carry_partial.len(),
266        });
267    }
268    reserve_bigint_output(final_sum, sum_partial.len(), "final_sum")?;
269    final_sum.clear();
270    let mut carry_in: u32 = 0;
271    for (sum, carry) in sum_partial.iter().zip(carry_partial.iter()) {
272        let (with_in, overflow_from_in) = sum.overflowing_add(carry_in);
273        final_sum.push(with_in);
274        // Total carry-out of this limb = original carry_partial OR
275        // carry-from-adding-the-incoming-carry. They cannot both fire
276        // at the same time unless sum was exactly 0xFFFF_FFFF (in which
277        // case the original add did NOT overflow; only the +1 did).
278        carry_in = *carry | u32::from(overflow_from_in);
279    }
280    Ok(carry_in)
281}
282
283#[cfg(any(test, feature = "cpu-parity"))]
284fn reserve_bigint_output(
285    out: &mut Vec<u32>,
286    len: usize,
287    operation: &'static str,
288) -> Result<(), BigIntAddCarryError> {
289    if len > out.capacity() {
290        crate::graph::scratch::reserve_graph_items(
291            out,
292            len - out.len(),
293            "bigint add-carry CPU oracle",
294            operation,
295        )
296        .map_err(|message| BigIntAddCarryError::AllocationFailed { operation, message })?;
297    }
298    Ok(())
299}
300
301#[cfg(feature = "inventory-registry")]
302inventory::submit! {
303    vyre_foundation::operation::OperationRegistration::primitive(
304        OP_ID,
305        || bigint_add_carry(4),
306        Some(|| {
307            vec![vec![
308                crate::wire::pack_u32_slice(&[1, u32::MAX, 5, u32::MAX]),
309                crate::wire::pack_u32_slice(&[2, 1, u32::MAX, u32::MAX]),
310                crate::wire::pack_u32_slice(&[0; 4]),
311                crate::wire::pack_u32_slice(&[0; 4]),
312            ]]
313        }),
314        Some(|| {
315            vec![vec![
316                crate::wire::pack_u32_slice(&[3, 0, 4, u32::MAX - 1]),
317                crate::wire::pack_u32_slice(&[0, 1, 1, 1]),
318            ]]
319        }),
320    )
321}
322
323#[cfg(test)]
324mod tests {
325    use super::*;
326
327    #[test]
328    fn cpu_zero_plus_zero_returns_zero_with_no_carries() {
329        let (sum, carry) =
330            bigint_add_carry_cpu(&[0, 0, 0, 0], &[0, 0, 0, 0]).expect("Fix: matching limbs");
331        assert_eq!(sum, vec![0, 0, 0, 0]);
332        assert_eq!(carry, vec![0, 0, 0, 0]);
333    }
334
335    #[test]
336    fn cpu_no_overflow_per_limb_keeps_carries_zero() {
337        let a = [1u32, 2, 3, 4];
338        let b = [10u32, 20, 30, 40];
339        let (sum, carry) = bigint_add_carry_cpu(&a, &b).expect("Fix: matching limbs");
340        assert_eq!(sum, vec![11, 22, 33, 44]);
341        assert_eq!(carry, vec![0, 0, 0, 0]);
342    }
343
344    #[test]
345    fn cpu_per_limb_overflow_emits_carry_bit() {
346        // limb 0: 0xFFFF_FFFF + 1 = 0, carry 1.
347        // limb 1: 0xFFFF_FFFF + 0 = 0xFFFF_FFFF, carry 0.
348        let a = [0xFFFF_FFFFu32, 0xFFFF_FFFFu32];
349        let b = [1u32, 0u32];
350        let (sum, carry) = bigint_add_carry_cpu(&a, &b).expect("Fix: matching limbs");
351        assert_eq!(sum, vec![0, 0xFFFF_FFFF]);
352        assert_eq!(carry, vec![1, 0]);
353    }
354
355    #[test]
356    fn cpu_max_plus_max_emits_per_limb_carry_and_truncated_sum() {
357        let a = [0xFFFF_FFFFu32; 4];
358        let b = [0xFFFF_FFFFu32; 4];
359        let (sum, carry) = bigint_add_carry_cpu(&a, &b).expect("Fix: matching limbs");
360        // 0xFFFF_FFFF + 0xFFFF_FFFF = 0x1_FFFF_FFFE (wraps to 0xFFFF_FFFE,
361        // carry 1) for every limb.
362        assert_eq!(sum, vec![0xFFFF_FFFEu32; 4]);
363        assert_eq!(carry, vec![1u32; 4]);
364    }
365
366    #[test]
367    fn resolve_carry_chain_propagates_single_carry_through_zeros() {
368        // sum_partial = [0xFFFF_FFFF, 0, 0, 0], carry_partial = [1, 0, 0, 0].
369        // After resolve: limb 0 stays 0xFFFF_FFFF; carry 1 propagates upward.
370        let sum_partial = vec![0xFFFF_FFFFu32, 0, 0, 0];
371        let carry_partial = vec![1u32, 0, 0, 0];
372        let (final_sum, final_carry) = resolve_carry_chain_cpu(&sum_partial, &carry_partial)
373            .expect("Fix: matching split limbs");
374        // limb 0 has no carry-in → stays 0xFFFF_FFFF.
375        assert_eq!(final_sum, vec![0xFFFF_FFFF, 1, 0, 0]);
376        assert_eq!(
377            final_carry, 0,
378            "the carry from limb 0 propagates into limb 1, then dies"
379        );
380    }
381
382    #[test]
383    fn resolve_carry_chain_handles_chained_overflow() {
384        // Adding 0x..FF + 0x..01 across all limbs ripples a carry the whole way.
385        // a = [0xFFFF_FFFF, 0xFFFF_FFFF, 0xFFFF_FFFF, 0]
386        // b = [0x0000_0001, 0x0000_0000, 0x0000_0000, 0]
387        // Expected final sum = [0, 0, 0, 1], final carry-out = 0.
388        let a = [0xFFFF_FFFFu32, 0xFFFF_FFFFu32, 0xFFFF_FFFFu32, 0];
389        let b = [1u32, 0, 0, 0];
390        let (sum_partial, carry_partial) =
391            bigint_add_carry_cpu(&a, &b).expect("Fix: matching limbs");
392        let (final_sum, final_carry) = resolve_carry_chain_cpu(&sum_partial, &carry_partial)
393            .expect("Fix: matching split limbs");
394        assert_eq!(final_sum, vec![0, 0, 0, 1]);
395        assert_eq!(final_carry, 0);
396    }
397
398    #[test]
399    fn resolve_carry_chain_emits_final_carry_out_at_top() {
400        // Adding the max two-limb integer to itself  -  the final carry-out
401        // must be 1 (the answer doesn't fit in 64 bits).
402        let a = [0xFFFF_FFFFu32, 0xFFFF_FFFFu32];
403        let b = [0xFFFF_FFFFu32, 0xFFFF_FFFFu32];
404        let (sum_partial, carry_partial) =
405            bigint_add_carry_cpu(&a, &b).expect("Fix: matching limbs");
406        let (_final_sum, final_carry) = resolve_carry_chain_cpu(&sum_partial, &carry_partial)
407            .expect("Fix: matching split limbs");
408        assert_eq!(
409            final_carry, 1,
410            "max + max in 64 bits overflows into the 65th bit"
411        );
412    }
413
414    #[test]
415    fn resolve_carry_chain_handles_corner_carry_in_only() {
416        // sum_partial = [0xFFFF_FFFF, 0xFFFF_FFFF], carry_partial = [1, 0].
417        // limb 0 → 0xFFFF_FFFF (no carry-in), then carry from limb 0 = 1.
418        // limb 1 → 0xFFFF_FFFF + 1 = 0, with overflow → next carry = 1.
419        let sum_partial = vec![0xFFFF_FFFFu32, 0xFFFF_FFFFu32];
420        let carry_partial = vec![1u32, 0];
421        let (final_sum, final_carry) = resolve_carry_chain_cpu(&sum_partial, &carry_partial)
422            .expect("Fix: matching split limbs");
423        assert_eq!(final_sum, vec![0xFFFF_FFFF, 0]);
424        assert_eq!(
425            final_carry, 1,
426            "carry propagated into limb 1 made it overflow"
427        );
428    }
429
430    #[test]
431    fn cpu_handles_8_limb_256_bit_operands() {
432        // 256-bit RSA-shape operand. Verifies the primitive scales to the
433        // sizes used by ECDSA / X25519.
434        let a = [0x1234_5678u32; 8];
435        let b = [0x8765_4321u32; 8];
436        let (sum, carry) = bigint_add_carry_cpu(&a, &b).expect("Fix: matching limbs");
437        // 0x1234_5678 + 0x8765_4321 = 0x9999_9999, no overflow.
438        assert_eq!(sum, vec![0x9999_9999u32; 8]);
439        assert_eq!(carry, vec![0u32; 8]);
440    }
441
442    #[test]
443    fn cpu_handles_128_limb_4096_bit_operands() {
444        // 4096-bit RSA modulus. Verifies the primitive scales to the
445        // sizes used by RSA-4096.
446        let a = vec![0x5555_5555u32; 128];
447        let b = vec![0xAAAA_AAAAu32; 128];
448        let (sum, carry) = bigint_add_carry_cpu(&a, &b).expect("Fix: matching limbs");
449        // 0x5555_5555 + 0xAAAA_AAAA = 0xFFFF_FFFF, no overflow.
450        assert_eq!(sum, vec![0xFFFF_FFFFu32; 128]);
451        assert_eq!(carry, vec![0u32; 128]);
452    }
453
454    #[test]
455    fn cpu_mismatched_limb_count_returns_error() {
456        let a = vec![0u32; 4];
457        let b = vec![0u32; 5];
458        assert_eq!(
459            bigint_add_carry_cpu(&a, &b),
460            Err(BigIntAddCarryError::LimbCountMismatch { a_len: 4, b_len: 5 })
461        );
462    }
463
464    #[test]
465    fn cpu_into_reuses_output_capacity() {
466        let a = [1u32, u32::MAX];
467        let b = [2u32, 1];
468
469        let mut sum = Vec::with_capacity(32);
470        let mut carry = Vec::with_capacity(32);
471        let sum_cap = sum.capacity();
472        let carry_cap = carry.capacity();
473        bigint_add_carry_cpu_into(&a, &b, &mut sum, &mut carry).expect("Fix: matching limbs");
474        assert_eq!(sum, vec![3, 0]);
475        assert_eq!(carry, vec![0, 1]);
476        assert_eq!(sum.capacity(), sum_cap);
477        assert_eq!(carry.capacity(), carry_cap);
478    }
479
480    #[test]
481    fn cpu_into_truncates_stale_tail_without_reallocating() {
482        let a = [1u32, u32::MAX];
483        let b = [2u32, 1];
484        let mut sum = Vec::with_capacity(8);
485        let mut carry = Vec::with_capacity(8);
486        sum.extend([99u32; 8]);
487        carry.extend([99u32; 8]);
488        let sum_ptr = sum.as_ptr();
489        let carry_ptr = carry.as_ptr();
490
491        bigint_add_carry_cpu_into(&a, &b, &mut sum, &mut carry).unwrap();
492
493        assert_eq!(sum, vec![3, 0]);
494        assert_eq!(carry, vec![0, 1]);
495        assert_eq!(sum.as_ptr(), sum_ptr);
496        assert_eq!(carry.as_ptr(), carry_ptr);
497    }
498
499    #[test]
500    fn resolve_into_truncates_stale_tail_without_reallocating() {
501        let mut out = Vec::with_capacity(8);
502        out.extend([99u32; 8]);
503        let ptr = out.as_ptr();
504
505        let carry = resolve_carry_chain_cpu_into(&[u32::MAX, u32::MAX], &[1, 0], &mut out).unwrap();
506
507        assert_eq!(out, vec![u32::MAX, 0]);
508        assert_eq!(carry, 1);
509        assert_eq!(out.as_ptr(), ptr);
510    }
511
512    #[test]
513    fn generated_split_and_resolve_matches_ripple_reference() {
514        let mut state = 0xB16A_DDCA_u32;
515        for case in 0..4096u32 {
516            state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
517            let len = match case {
518                0 => 1,
519                1 => 24,
520                2 => 256,
521                3 => 257,
522                4 => 1025,
523                _ => state % 4097 + 1,
524            } as usize;
525            let mut a = Vec::with_capacity(len);
526            let mut b = Vec::with_capacity(len);
527            for idx in 0..len {
528                state = state.rotate_left(9) ^ (idx as u32).wrapping_mul(0x9E37_79B9);
529                let left = match idx % 11 {
530                    0 => u32::MAX,
531                    1 => 0,
532                    2 => 0x8000_0000,
533                    _ => state,
534                };
535                let right = match idx % 13 {
536                    0 => 1,
537                    1 => u32::MAX,
538                    2 => 0x8000_0000,
539                    _ => state.rotate_right(7),
540                };
541                a.push(left);
542                b.push(right);
543            }
544            let (sum_partial, carry_partial) = bigint_add_carry_cpu(&a, &b).unwrap();
545            let (final_sum, final_carry) =
546                resolve_carry_chain_cpu(&sum_partial, &carry_partial).unwrap();
547            let mut expected = Vec::with_capacity(len);
548            let mut carry = 0u64;
549            for i in 0..len {
550                let total = a[i] as u64 + b[i] as u64 + carry;
551                expected.push(total as u32);
552                carry = total >> 32;
553            }
554
555            assert_eq!(
556                final_sum, expected,
557                "generated bigint case {case} len={len}"
558            );
559            assert_eq!(
560                final_carry, carry as u32,
561                "generated bigint carry case {case} len={len}"
562            );
563        }
564    }
565
566    #[test]
567    fn resolve_carry_chain_rejects_length_mismatch() {
568        let mut out = Vec::new();
569        assert_eq!(
570            resolve_carry_chain_cpu_into(&[0, 1], &[0], &mut out),
571            Err(BigIntAddCarryError::SplitCarryLengthMismatch {
572                sum_len: 2,
573                carry_len: 1,
574            })
575        );
576    }
577
578    #[test]
579    fn build_program_returns_well_formed_program() {
580        let program = bigint_add_carry(8);
581        assert_eq!(
582            program.buffers().len(),
583            4,
584            "a, b, sum_partial, carry_partial"
585        );
586        assert_eq!(program.workgroup_size(), BIGINT_ADD_CARRY_WORKGROUP_SIZE);
587    }
588
589    #[test]
590    fn dispatch_grid_packs_limb_lanes_into_workgroups() {
591        assert_eq!(bigint_add_carry_dispatch_grid(0), [1, 1, 1]);
592        assert_eq!(bigint_add_carry_dispatch_grid(1), [1, 1, 1]);
593        assert_eq!(bigint_add_carry_dispatch_grid(256), [1, 1, 1]);
594        assert_eq!(bigint_add_carry_dispatch_grid(257), [2, 1, 1]);
595        assert_eq!(bigint_add_carry_dispatch_grid(1025), [5, 1, 1]);
596    }
597
598    #[test]
599    fn zero_limb_count_traps() {
600        let program = bigint_add_carry(0);
601        assert!(program.stats().trap());
602    }
603
604    #[test]
605    fn build_program_is_deterministic_across_calls() {
606        // Same input → same Program. This is the wire-content-hash
607        // contract; if it ever fails, differential compilation breaks.
608        let p1 = bigint_add_carry(16);
609        let p2 = bigint_add_carry(16);
610        assert_eq!(
611            p1.buffers().len(),
612            p2.buffers().len(),
613            "two builds with identical inputs must produce identical buffer lists"
614        );
615        assert_eq!(p1.workgroup_size(), p2.workgroup_size());
616    }
617
618    #[test]
619    fn op_id_is_canonical_and_stable() {
620        // Op ids are wire-format-visible; changing them is a breaking change.
621        assert_eq!(OP_ID, "vyre-primitives::math::bigint_add_carry");
622    }
623
624    #[test]
625    fn binding_indices_are_canonical_and_stable() {
626        // Bindings are wire-format-visible; changing them is a breaking change.
627        assert_eq!(BINDING_A_IN, 0);
628        assert_eq!(BINDING_B_IN, 1);
629        assert_eq!(BINDING_SUM_PARTIAL_OUT, 2);
630        assert_eq!(BINDING_CARRY_PARTIAL_OUT, 3);
631    }
632}