code wiki / (root) / nx_u256_mul_test.nx

nx_u256_mul_test.nx source

↩ module page · 207 lines · 5854 B

1// nx_u256_mul_test.nx -- KAT for the 256x256 -> 512 wide-multiply. 2// 3// Verifies: 4// - 0 * 0 = 0 5// - 1 * 1 = 1 6// - a * 1 = a (identity) 7// - 1 * a = a (commutative-with-identity) 8// - 2^32 * 2^32 = 2^64 (cross-limb single bit) 9// - (2^256 - 1)^2 produces the well-known wide value 10// - small * small produces known products (3 * 5 = 15, etc.) 11// - large limb-spanning case (0xFFFFFFFF * 0xFFFFFFFF) 12// - wide_cmp / wide_fits_in_256 / wide_copy_low helpers 13// 14// expect_exit: 0 15// license_tier: ORIGINAL 16 17import "nx_syscalls.nx" 18import "nx_u256.nx" 19import "nx_u256_mul.nx" 20 21func main() -> i64 { 22 let a: *i64 = u256_alloc() 23 let b: *i64 = u256_alloc() 24 let out: *i64 = u256_wide_alloc() 25 let tmp: *i64 = u256_wide_alloc() 26 27 // ---- Test A: 0 * 0 = 0 ---- 28 u256_zero(a); u256_zero(b) 29 u256_mul_wide(out, a, b) 30 var i: i64 = 0 31 while i < 16 { 32 if out[i] != 0 { return 1 } 33 i = i + 1 34 } 35 36 // ---- Test B: 1 * 1 = 1 ---- 37 u256_one(a); u256_one(b) 38 u256_mul_wide(out, a, b) 39 if out[0] != 1 { return 2 } 40 i = 1 41 while i < 16 { 42 if out[i] != 0 { return 3 } 43 i = i + 1 44 } 45 46 // ---- Test C: 3 * 5 = 15 ---- 47 u256_zero(a); a[0] = 3 48 u256_zero(b); b[0] = 5 49 u256_mul_wide(out, a, b) 50 if out[0] != 15 { return 4 } 51 i = 1 52 while i < 16 { 53 if out[i] != 0 { return 5 } 54 i = i + 1 55 } 56 57 // ---- Test D: 0xFFFFFFFF * 0xFFFFFFFF = 0xFFFFFFFE00000001 ---- 58 // (well-known: (2^32 - 1)^2 = 2^64 - 2^33 + 1) 59 u256_zero(a); a[0] = 0xFFFFFFFF 60 u256_zero(b); b[0] = 0xFFFFFFFF 61 u256_mul_wide(out, a, b) 62 if out[0] != 0x00000001 { return 6 } 63 if out[1] != 0xFFFFFFFE { return 7 } 64 i = 2 65 while i < 16 { 66 if out[i] != 0 { return 8 } 67 i = i + 1 68 } 69 70 // ---- Test E: 2^32 * 2^32 = 2^64 ---- 71 // (Bit cross-limb: limb[1]=1) * (limb[1]=1) -> limb[2]=1 72 u256_zero(a); a[1] = 1 73 u256_zero(b); b[1] = 1 74 u256_mul_wide(out, a, b) 75 if out[0] != 0 { return 10 } 76 if out[1] != 0 { return 11 } 77 if out[2] != 1 { return 12 } 78 i = 3 79 while i < 16 { 80 if out[i] != 0 { return 13 } 81 i = i + 1 82 } 83 84 // ---- Test F: 2^255 * 2 = 2^256 ---- 85 // 2^255: limb[7] = 0x80000000 86 // 2: limb[0] = 2 87 // product: 2^256 -> low half all zero, high half[0] = 1 88 u256_zero(a); a[7] = 0x80000000 89 u256_zero(b); b[0] = 2 90 u256_mul_wide(out, a, b) 91 if out[0] != 0 { return 20 } 92 if out[7] != 0 { return 21 } 93 if out[8] != 1 { return 22 } 94 i = 9 95 while i < 16 { 96 if out[i] != 0 { return 23 } 97 i = i + 1 98 } 99 100 // ---- Test G: (2^256 - 1) * 1 = 2^256 - 1 ---- 101 // a = all 0xFFFFFFFF (max u256), b = 1 102 i = 0 103 while i < 8 { a[i] = 0xFFFFFFFF; i = i + 1 } 104 u256_one(b) 105 u256_mul_wide(out, a, b) 106 i = 0 107 while i < 8 { 108 if out[i] != 0xFFFFFFFF { return 30 + i } 109 i = i + 1 110 } 111 i = 8 112 while i < 16 { 113 if out[i] != 0 { return 40 + i } 114 i = i + 1 115 } 116 117 // ---- Test H: (2^256 - 1)^2 ---- 118 // Result is (2^256 - 1)^2 = 2^512 - 2^257 + 1 119 // In 16-limb form (LE): 120 // limb[0] = 1 121 // limb[1..7] = 0 122 // limb[8..15] each = ? 123 // Compute via mathematical formula: 124 // 2^512 - 2*2^256 + 1 125 // limb[0] = 1 126 // limb[8] = (0 - 2) mod 2^32 = 0xFFFFFFFE plus borrow 127 // limbs[9..15] are propagation of the -2 borrow 128 // 129 // Simpler verification: just check that the result is bit- 130 // exact against (max * max) by independently computing 131 // a*max via repeated subtraction or by checking specific 132 // limbs we know. 133 // 134 // From the schoolbook expansion of (2^256-1)^2: 135 // low limb: 1 136 // limbs 1..7: 0 137 // limb 8: 0xFFFFFFFE 138 // limbs 9..15: 0xFFFFFFFF 139 i = 0 140 while i < 8 { a[i] = 0xFFFFFFFF; i = i + 1 } 141 i = 0 142 while i < 8 { b[i] = 0xFFFFFFFF; i = i + 1 } 143 u256_mul_wide(out, a, b) 144 if out[0] != 1 { return 50 } 145 i = 1 146 while i < 8 { 147 if out[i] != 0 { return 51 } 148 i = i + 1 149 } 150 if out[8] != 0xFFFFFFFE { return 52 } 151 i = 9 152 while i < 16 { 153 if out[i] != 0xFFFFFFFF { return 53 } 154 i = i + 1 155 } 156 157 // ---- Test I: a * 1 = a for various a ---- 158 let bytes: *u8 = sys_mmap(32) 159 var k: i64 = 0 160 while k < 32 { bytes[k] = (0x10 + k) as u8; k = k + 1 } 161 u256_load_be(a, bytes) 162 u256_one(b) 163 u256_mul_wide(out, a, b) 164 // out[0..8] must equal a, out[8..16] must be 0 165 i = 0 166 while i < 8 { 167 if (out[i] & 0xFFFFFFFF) != (a[i] & 0xFFFFFFFF) { return 60 + i } 168 i = i + 1 169 } 170 if u256_wide_fits_in_256(out) != 1 { return 70 } 171 172 // ---- Test J: 1 * a = a (verify commutative-with-identity) ---- 173 u256_mul_wide(tmp, b, a) // 1 * a 174 if u256_wide_cmp(tmp, out) != 0 { return 80 } 175 176 // ---- Test K: wide_copy_low extracts low 256 bits ---- 177 let low: *i64 = u256_alloc() 178 u256_wide_copy_low(low, out) // out is currently a*1 179 if u256_eq(low, a) != 1 { return 90 } 180 181 // ---- Test L: wide_fits_in_256 boundary ---- 182 // Build a wide value with upper limb non-zero -> should return 0 183 u256_one(b) 184 i = 0 185 while i < 16 { tmp[i] = 0; i = i + 1 } 186 tmp[8] = 1 187 if u256_wide_fits_in_256(tmp) != 0 { return 100 } 188 tmp[8] = 0 189 tmp[15] = 1 190 if u256_wide_fits_in_256(tmp) != 0 { return 101 } 191 tmp[15] = 0 192 if u256_wide_fits_in_256(tmp) != 1 { return 102 } 193 194 // ---- Test M: wide_cmp ---- 195 i = 0 196 while i < 16 { out[i] = 0; tmp[i] = 0; i = i + 1 } 197 out[0] = 5 198 tmp[0] = 7 199 if u256_wide_cmp(out, tmp) != (0 - 1) { return 110 } 200 if u256_wide_cmp(tmp, out) != 1 { return 111 } 201 out[0] = 7 202 if u256_wide_cmp(out, tmp) != 0 { return 112 } 203 out[15] = 1 204 if u256_wide_cmp(out, tmp) != 1 { return 113 } 205 206 return 0 207}