code wiki / (root) / nx_demodulation_test.nx

nx_demodulation_test.nx source

↩ module page · 181 lines · 7358 B

1// nx_demodulation_test.nx -- demodulation smoke. 2 3import "nx_syscalls.nx" 4import "nx_runtime.nx" 5import "nx_tier.nx" 6import "nx_result.nx" 7import "nx_unify.nx" 8import "nx_resolution.nx" 9import "nx_term_order.nx" 10import "nx_demodulation.nx" 11 12const SYM_A: nx_int = 100 13const SYM_B: nx_int = 101 14const SYM_F: nx_int = 200 15const SYM_G: nx_int = 201 16const SYM_H: nx_int = 250 17 18const VAR_X: nx_int = 0 19const VAR_Y: nx_int = 1 20const NX_MAX_VAR: nx_int = 4 21 22func mk_unary(sym: nx_int, child: *Term) -> *Term { 23 let buf: *Term = (sys_mmap(NX_TERM_BYTES as i64)) as *Term 24 buf.kind = child.kind; buf.sym = child.sym 25 buf.n_args = child.n_args; buf.args = child.args 26 return nx_term_app(sym, 1, buf) 27} 28 29func mk_binary(sym: nx_int, c0: *Term, c1: *Term) -> *Term { 30 let args: *Term = (sys_mmap((2 * NX_TERM_BYTES) as i64)) as *Term 31 let a0: *Term = args 32 a0.kind = c0.kind; a0.sym = c0.sym; a0.n_args = c0.n_args; a0.args = c0.args 33 let a1: *Term = ((args as nx_int) + NX_TERM_BYTES) as *Term 34 a1.kind = c1.kind; a1.sym = c1.sym; a1.n_args = c1.n_args; a1.args = c1.args 35 return nx_term_app(sym, 2, args) 36} 37 38// Build an admissible KBO state: A < B < F < G < H in precedence, 39// all weights = 1, w_0 = 1. 40func mk_kbo() -> *KboState { 41 let st: *KboState = nx_kbo_new(1) 42 let _r1: *NxResult = nx_kbo_register(st, SYM_A, 1, 10, 0) 43 let _r2: *NxResult = nx_kbo_register(st, SYM_B, 1, 11, 0) 44 let _r3: *NxResult = nx_kbo_register(st, SYM_F, 1, 20, 1) 45 let _r4: *NxResult = nx_kbo_register(st, SYM_G, 1, 21, 1) 46 let _r5: *NxResult = nx_kbo_register(st, SYM_H, 1, 30, 2) 47 return st 48} 49 50func report(name: *u8, expected_sym: nx_int, actual: *Term) -> nx_int { 51 print(" " as *u8); print(name); print(" -> head=" as *u8); print_i64(actual.sym) 52 if actual.sym == expected_sym { println(" PASS" as *u8); return 0 } 53 print(" (expected " as *u8); print_i64(expected_sym); println(") FAIL" as *u8) 54 return 1 55} 56 57func main() -> nx_exit { 58 println("=== Demodulation smoke ===" as *u8) 59 60 let st: *KboState = mk_kbo() 61 var fails: nx_int = 0 62 63 // ---------- Test 1: orient an equation ---------------------- 64 // Equation: a = b (B has higher precedence -> b > a; lhs should be b) 65 let r1: *NxResult = nx_eqn_orient(st, nx_term_const(SYM_A), nx_term_const(SYM_B), NX_MAX_VAR) 66 if nx_result_is_err(r1) == 1 { 67 println("1. orient(a, b) FAIL: orient returned ERR" as *u8) 68 fails = fails + 1 69 } else { 70 let eqn1: *Equation = (nx_result_unwrap(r1)) as *Equation 71 if eqn1.lhs.sym == SYM_B { 72 println("1. orient(a, b) -> lhs=b, rhs=a PASS" as *u8) 73 } else { 74 print("1. orient(a, b) -> lhs.sym=" as *u8); print_i64(eqn1.lhs.sym); println(" FAIL" as *u8) 75 fails = fails + 1 76 } 77 } 78 79 // ---------- Test 2: no-match leaves term unchanged ---------- 80 // Equation: g(x) -> f(x); target: a (no g in target) 81 // Since g > f in precedence and equal weight, g(x) > f(x). 82 let r2: *NxResult = nx_eqn_orient(st, mk_unary(SYM_G, nx_term_var(VAR_X)), 83 mk_unary(SYM_F, nx_term_var(VAR_X)), NX_MAX_VAR) 84 let eqn2: *Equation = (nx_result_unwrap(r2)) as *Equation 85 let t2: *Term = nx_term_const(SYM_A) 86 let out2: *Term = nx_demodulate_term(eqn2, t2) 87 if (out2 as nx_int) == (t2 as nx_int) { 88 println("2. demod(g(x)->f(x), a) -> unchanged PASS" as *u8) 89 } else { 90 println("2. demod(g(x)->f(x), a) -> CHANGED FAIL" as *u8) 91 fails = fails + 1 92 } 93 94 // ---------- Test 3: root-position rewrite ------------------- 95 // Equation: g(x) -> f(x); target: g(a) --> f(a) 96 let t3: *Term = mk_unary(SYM_G, nx_term_const(SYM_A)) 97 let out3: *Term = nx_demodulate_term(eqn2, t3) 98 fails = fails + report("3. demod(g(x)->f(x), g(a))" as *u8, SYM_F, out3) 99 100 // ---------- Test 4: subterm rewrite ------------------------- 101 // Equation: g(x) -> f(x); target: f(g(a)) --> f(f(a)) 102 let t4_inner: *Term = mk_unary(SYM_G, nx_term_const(SYM_A)) 103 let t4: *Term = mk_unary(SYM_F, t4_inner) 104 let out4: *Term = nx_demodulate_term(eqn2, t4) 105 let out4_child: *Term = nx_term_arg(out4, 0) 106 if out4.sym == SYM_F { 107 if out4_child.sym == SYM_F { 108 println("4. demod(g(x)->f(x), f(g(a))) -> f(f(a)) PASS" as *u8) 109 } else { 110 print("4. demod(g(x)->f(x), f(g(a))) -> child.sym=" as *u8) 111 print_i64(out4_child.sym); println(" FAIL" as *u8) 112 fails = fails + 1 113 } 114 } else { 115 print("4. demod(g(x)->f(x), f(g(a))) -> head=" as *u8) 116 print_i64(out4.sym); println(" FAIL" as *u8) 117 fails = fails + 1 118 } 119 120 // ---------- Test 5: multiple rewrites in one term ------------ 121 // Equation: g(x) -> f(x); target: h(g(a), g(b)) --> h(f(a), f(b)) 122 let t5: *Term = mk_binary(SYM_H, 123 mk_unary(SYM_G, nx_term_const(SYM_A)), 124 mk_unary(SYM_G, nx_term_const(SYM_B))) 125 let out5: *Term = nx_demodulate_term(eqn2, t5) 126 let out5_l: *Term = nx_term_arg(out5, 0) 127 let out5_r: *Term = nx_term_arg(out5, 1) 128 if out5_l.sym == SYM_F { 129 if out5_r.sym == SYM_F { 130 println("5. demod(g(x)->f(x), h(g(a),g(b))) -> h(f(a),f(b)) PASS" as *u8) 131 } else { 132 println("5. demod 5 right-side mismatch FAIL" as *u8); fails = fails + 1 133 } 134 } else { 135 println("5. demod 5 left-side mismatch FAIL" as *u8); fails = fails + 1 136 } 137 138 // ---------- Test 6: clause-level demodulation ---------------- 139 // Equation: g(x) -> f(x); clause: { p(g(a)), p(b) } 140 // Expected: { p(f(a)), p(b) }, n_changed = 1 141 let SYM_P: nx_int = 300 142 let _rp: *NxResult = nx_kbo_register(st, SYM_P, 1, 5, 1) 143 let cin: *Clause = nx_clause_new() 144 let _r6a: *NxResult = nx_clause_add(cin, nx_lit_make(NX_LIT_POS, 145 mk_unary(SYM_P, mk_unary(SYM_G, nx_term_const(SYM_A))))) 146 let _r6b: *NxResult = nx_clause_add(cin, nx_lit_make(NX_LIT_POS, 147 mk_unary(SYM_P, nx_term_const(SYM_B)))) 148 let cout: *Clause = nx_clause_new() 149 let n_changed: nx_int = nx_demodulate_clause(eqn2, cin, cout) 150 if n_changed == 1 { 151 if cout.n_lits == 2 { 152 println("6. demod_clause -> 1 lit changed, 2 lits total PASS" as *u8) 153 } else { 154 print("6. demod_clause -> n_lits=" as *u8); print_i64(cout.n_lits); println(" FAIL" as *u8) 155 fails = fails + 1 156 } 157 } else { 158 print("6. demod_clause -> n_changed=" as *u8); print_i64(n_changed); println(" FAIL" as *u8) 159 fails = fails + 1 160 } 161 162 // ---------- Test 7: orient refuses incomparable equation ------ 163 // f(x) and g(y) -- different vars, incomparable. Must reject. 164 let r7: *NxResult = nx_eqn_orient(st, 165 mk_unary(SYM_F, nx_term_var(VAR_X)), 166 mk_unary(SYM_G, nx_term_var(VAR_Y)), NX_MAX_VAR) 167 if nx_result_is_err(r7) == 1 { 168 println("7. orient(f(x), g(y)) -> rejected (incomparable) PASS" as *u8) 169 } else { 170 println("7. orient(f(x), g(y)) -> ACCEPTED FAIL" as *u8) 171 fails = fails + 1 172 } 173 174 println("" as *u8) 175 if fails == 0 { 176 println("=== ALL 7 demodulation tests PASS ===" as *u8) 177 return 0 178 } 179 print("=== " as *u8); print_i64(fails); println(" demodulation tests FAILED ===" as *u8) 180 return 1 181}