code wiki / (root) / nx_p256_field_mul_oracle_test.nx

nx_p256_field_mul_oracle_test.nx source

↩ module page · 96 lines · 2967 B

1// nx_p256_field_mul_oracle_test.nx -- validates the FAST Solinas 2// p256_field_mul against the SLOW bit-serial oracle (p256_field_mul_slow) 3// over many random inputs + edge cases. fast MUST equal slow for every 4// canonical (a,b). This is the correctness gate for the Solinas path. 5 6import "nx_syscalls.nx" 7import "nx_u256.nx" 8import "nx_csprng.nx" 9import "nx_p256_field.nx" 10import "nx_p256_field_mul.nx" 11 12// a >= p ? (1 yes, 0 no) 13func _t_ge_p(a: *i64, p: *i64) -> i64 { 14 var i: i64 = NX_U256_LIMBS - 1 15 while i >= 0 { 16 let av: i64 = a[i] & NX_U256_LIMB_MASK 17 let pv: i64 = p[i] & NX_U256_LIMB_MASK 18 if av > pv { return 1 } 19 if av < pv { return 0 } 20 i = i - 1 21 } 22 return 1 23} 24 25// a -= p (8-limb, assumes a >= p) 26func _t_sub_p(a: *i64, p: *i64) -> i64 { 27 var borrow: i64 = 0 28 var j: i64 = 0 29 while j < NX_U256_LIMBS { 30 let d: i64 = (a[j] & NX_U256_LIMB_MASK) - (p[j] & NX_U256_LIMB_MASK) - borrow 31 if d < 0 { a[j] = (d + (1 << NX_U256_LIMB_BITS)) & NX_U256_LIMB_MASK; borrow = 1 } 32 else { a[j] = d & NX_U256_LIMB_MASK; borrow = 0 } 33 j = j + 1 34 } 35 return 0 36} 37 38func _t_canon(a: *i64, p: *i64) -> i64 { 39 if _t_ge_p(a, p) == 1 { _t_sub_p(a, p) } 40 return 0 41} 42 43func _t_check(a: *i64, b: *i64) -> i64 { 44 let fast: *i64 = u256_alloc() 45 let slow: *i64 = u256_alloc() 46 p256_field_mul(fast, a, b) 47 p256_field_mul_slow(slow, a, b) 48 if u256_eq(fast, slow) == 1 { return 0 } 49 return 1 50} 51 52func main() -> i64 { 53 let p: *i64 = u256_alloc() 54 p256_field_load_p(p) 55 let a: *i64 = u256_alloc() 56 let b: *i64 = u256_alloc() 57 let rnd: *u8 = sys_mmap(32) 58 var fails: i64 = 0 59 60 // Random canonical inputs. 61 var n: i64 = 0 62 while n < 600 { 63 nx_csprng_fill(rnd, 32); u256_load_be(a, rnd); _t_canon(a, p) 64 nx_csprng_fill(rnd, 32); u256_load_be(b, rnd); _t_canon(b, p) 65 fails = fails + _t_check(a, b) 66 n = n + 1 67 } 68 69 // Edge cases. 70 // p-1 (construct: load p, subtract 1). 71 let pm1: *i64 = u256_alloc() 72 u256_copy(pm1, p) 73 var borrow: i64 = 1 74 var j: i64 = 0 75 while j < NX_U256_LIMBS { 76 let d: i64 = (pm1[j] & NX_U256_LIMB_MASK) - borrow 77 if d < 0 { pm1[j] = (d + (1 << NX_U256_LIMB_BITS)) & NX_U256_LIMB_MASK; borrow = 1 } 78 else { pm1[j] = d & NX_U256_LIMB_MASK; borrow = 0 } 79 j = j + 1 80 } 81 let zero: *i64 = u256_alloc(); u256_zero(zero) 82 let one: *i64 = u256_alloc(); u256_one(one) 83 fails = fails + _t_check(pm1, pm1) // (p-1)^2 -- max product 84 fails = fails + _t_check(pm1, one) 85 fails = fails + _t_check(pm1, zero) 86 fails = fails + _t_check(zero, zero) 87 fails = fails + _t_check(one, one) 88 fails = fails + _t_check(pm1, p) // p is not canonical but exercises high bits 89 90 if fails == 0 { 91 sys_write(1, "ORACLE PASS (606 cases)\n" as *u8, 24) 92 return 0 93 } 94 sys_write(1, "ORACLE FAIL\n" as *u8, 12) 95 return 1 96}