nx_p256_modn_mont_test.nx source
↩ module page · 81 lines · 2501 B
1// nx_p256_modn_mont_test.nx -- oracle test for the Montgomery fast
2// p256_modn_inv vs the slow Fermat reference, plus the a*a^-1 == 1
3// identity. Run on the NAS (x86_64). Expect: all PASS, 0 mismatches.
4import "nx_syscalls.nx"
5import "nx_csprng.nx"
6import "nx_p256_modn.nx"
7
8func t_strlen(s: *u8) -> i64 { var n: i64 = 0; while s[n] != 0 { n = n + 1 } return n }
9
10func t_print_num(label: *u8, v: i64) -> i64 {
11 sys_write(1, label, t_strlen(label))
12 let buf: *u8 = sys_mmap(24)
13 var d: i64 = 0
14 if v == 0 { buf[0] = 48; d = 1 } else {
15 var x: i64 = v
16 var c: i64 = 0
17 while x > 0 { c = c + 1; x = x / 10 }
18 d = c
19 var i: i64 = d - 1
20 x = v
21 while i >= 0 { buf[i] = (48 + (x % 10)) as u8; x = x / 10; i = i - 1 }
22 }
23 sys_write(1, buf, d)
24 sys_write(1, "\n" as *u8, 1)
25 return 0
26}
27
28// Set k to a random scalar in [1, n).
29func rand_scalar(k: *i64) -> i64 {
30 let be: *u8 = sys_mmap(32)
31 nx_csprng_fill(be, 32)
32 u256_load_be(k, be)
33 p256_modn_reduce(k, k)
34 if u256_is_zero(k) == 1 { p256_modn_one(k) }
35 return 0
36}
37
38func check_one(k: *i64, mismatch_box: *i64, identity_fail_box: *i64) -> i64 {
39 let fast: *i64 = u256_alloc()
40 let slow: *i64 = u256_alloc()
41 p256_modn_inv(fast, k)
42 p256_modn_inv_slow(slow, k)
43 if u256_eq(fast, slow) == 0 { mismatch_box[0] = mismatch_box[0] + 1 }
44 // identity: fast * k mod n == 1 (uses the slow, known-correct mul)
45 let chk: *i64 = u256_alloc()
46 p256_modn_mul(chk, fast, k)
47 let one: *i64 = u256_alloc()
48 p256_modn_one(one)
49 if u256_eq(chk, one) == 0 { identity_fail_box[0] = identity_fail_box[0] + 1 }
50 return 0
51}
52
53func main() -> i64 {
54 let mism: *i64 = (sys_mmap(8)) as *i64
55 let idf: *i64 = (sys_mmap(8)) as *i64
56 mism[0] = 0
57 idf[0] = 0
58
59 // Fixed edges: 1, 2, 3.
60 let k: *i64 = u256_alloc()
61 p256_modn_one(k); check_one(k, mism, idf)
62 u256_zero(k); k[0] = 2; check_one(k, mism, idf)
63 u256_zero(k); k[0] = 3; check_one(k, mism, idf)
64
65 // Random scalars.
66 let N: i64 = 100
67 var i: i64 = 0
68 while i < N {
69 rand_scalar(k)
70 check_one(k, mism, idf)
71 i = i + 1
72 }
73
74 t_print_num("cases =" as *u8, N + 3)
75 t_print_num("fast!=slow =" as *u8, mism[0])
76 t_print_num("inv*k!=1 =" as *u8, idf[0])
77 if mism[0] == 0 { if idf[0] == 0 {
78 sys_write(1, "RESULT: ALL PASS\n" as *u8, 17)
79 } }
80 return 0
81}