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}