code wiki / bin / nx_p256_comb_test.nx

nx_p256_comb_test.nx source

↩ module page · 68 lines · 2187 B

1// nx_p256_comb_test.nx -- oracle test: fixed-base comb k*G must equal the 2// generic double-and-add p256_scalar_mul(k, G) for all scalars. 3import "nx_syscalls.nx" 4import "nx_csprng.nx" 5import "nx_p256_comb.nx" 6import "nx_p256_scalar_mul.nx" 7 8func t_strlen(s: *u8) -> i64 { var n: i64 = 0; while s[n] != 0 { n = n + 1 } return n } 9func t_print_num(label: *u8, v: i64) -> i64 { 10 sys_write(1, label, t_strlen(label)) 11 let buf: *u8 = sys_mmap(24) 12 var d: i64 = 0 13 if v == 0 { buf[0] = 48; d = 1 } else { 14 var x: i64 = v 15 var c: i64 = 0 16 while x > 0 { c = c + 1; x = x / 10 } 17 d = c 18 var i: i64 = d - 1 19 x = v 20 while i >= 0 { buf[i] = (48 + (x % 10)) as u8; x = x / 10; i = i - 1 } 21 } 22 sys_write(1, buf, d) 23 sys_write(1, "\n" as *u8, 1) 24 return 0 25} 26 27func check_k(k: *i64, g: *P256Point, table: *i64, box: *i64) -> i64 { 28 let r1: *P256Point = p256_point_alloc() 29 let r2: *P256Point = p256_point_alloc() 30 p256_scalar_mul_base(r1, k, table) 31 p256_scalar_mul(r2, k, g) 32 if p256_point_eq(r1, r2) == 0 { box[0] = box[0] + 1 } 33 return 0 34} 35 36func main() -> i64 { 37 let table: *i64 = (sys_mmap(NX_P256_COMB_BYTES)) as *i64 38 p256_comb_build(table) 39 let g: *P256Point = p256_point_alloc() 40 p256_point_load_g(g) 41 let box: *i64 = (sys_mmap(8)) as *i64 42 box[0] = 0 43 let k: *i64 = u256_alloc() 44 45 // Edges: 0, 1, 2, 15, 16, 17. 46 u256_zero(k); check_k(k, g, table, box) 47 u256_zero(k); k[0] = 1; check_k(k, g, table, box) 48 u256_zero(k); k[0] = 2; check_k(k, g, table, box) 49 u256_zero(k); k[0] = 15; check_k(k, g, table, box) 50 u256_zero(k); k[0] = 16; check_k(k, g, table, box) 51 u256_zero(k); k[0] = 17; check_k(k, g, table, box) 52 53 // Random full-width scalars. 54 let N: i64 = 50 55 var i: i64 = 0 56 while i < N { 57 let be: *u8 = sys_mmap(32) 58 nx_csprng_fill(be, 32) 59 u256_load_be(k, be) 60 check_k(k, g, table, box) 61 i = i + 1 62 } 63 64 t_print_num("cases =" as *u8, N + 6) 65 t_print_num("comb!=dbladd=" as *u8, box[0]) 66 if box[0] == 0 { sys_write(1, "RESULT: ALL PASS\n" as *u8, 17) } 67 return 0 68}