1use std::sync::Arc;
44
45use vyre_foundation::ir::model::expr::Ident;
46use vyre_foundation::ir::{BufferAccess, BufferDecl, DataType, Expr, Node, Program};
47
48pub const OP_ID: &str = "vyre-primitives::math::bigint_add_carry";
50
51pub const BINDING_A_IN: u32 = 0;
53pub const BINDING_B_IN: u32 = 1;
55pub const BINDING_SUM_PARTIAL_OUT: u32 = 2;
57pub const BINDING_CARRY_PARTIAL_OUT: u32 = 3;
59
60pub const BIGINT_ADD_CARRY_WORKGROUP_SIZE: [u32; 3] = [256, 1, 1];
62
63#[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#[derive(Debug, Clone, PartialEq, Eq)]
79#[non_exhaustive]
80pub enum BigIntAddCarryError {
81 LimbCountMismatch {
83 a_len: usize,
85 b_len: usize,
87 },
88 SplitCarryLengthMismatch {
90 sum_len: usize,
92 carry_len: usize,
94 },
95 AllocationFailed {
97 operation: &'static str,
99 message: String,
101 },
102}
103
104#[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 Node::let_bind("sum", Expr::add(Expr::var("a_limb"), Expr::var("b_limb"))),
134 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#[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#[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#[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#[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 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 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 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 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 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 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 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 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 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 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 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 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 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 assert_eq!(OP_ID, "vyre-primitives::math::bigint_add_carry");
622 }
623
624 #[test]
625 fn binding_indices_are_canonical_and_stable() {
626 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}