cdclkit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
cdclkit/pyeq.py ADDED
@@ -0,0 +1,824 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # Copyright (c) 2026 Carlo Perassi. Licensed under the Apache License 2.0.
3
+ """Prove two Python functions agree on every input, or produce one that breaks.
4
+
5
+ >>> def slow(a, b): return a * 2 + b * 2
6
+ >>> def fast(a, b): return (a + b) << 1
7
+ >>> equivalent(slow, fast, widths={"a": 8, "b": 8}).proved
8
+ True
9
+
10
+ The question this answers is "did my refactor change behaviour?", and it
11
+ answers it by *proof* rather than by sampling. Tests check the inputs you
12
+ thought of. This checks all of them.
13
+
14
+ How
15
+ ---
16
+ Each function is compiled from its Python AST into a fixed-width bit-vector
17
+ circuit, the circuit into CNF via `cdclkit.encodings.Encoder`, and the two
18
+ circuits into a **miter**: shared inputs, outputs compared, the comparison
19
+ asserted to differ. Unsatisfiable means no input distinguishes them. Satisfiable
20
+ means the model *is* a distinguishing input, which is handed back.
21
+
22
+ The solving half of this already existed -- `examples/equivalence.py` miters two
23
+ adder implementations, and `bench/run_bench.py` builds an array multiplier out
24
+ of the same gates. This module is the front end.
25
+
26
+ **Semantics, and this matters**
27
+ -------------------------------
28
+ Python integers are arbitrary precision. This compiles them to **fixed-width
29
+ two's-complement with wrapping**, at a width you declare. So:
30
+
31
+ * "equivalent" means **equivalent as fixed-width machine integers**, not as
32
+ Python programs. Two functions that agree at 8 bits can differ in Python if
33
+ either exceeds the range, and vice versa;
34
+ * a counterexample is always real *at that width* -- the circuit outputs are
35
+ read back from the model and asserted to differ, so a spurious one is a bug
36
+ here and says so;
37
+ * when the circuits differ but Python agrees, the result sets
38
+ ``overflow_only`` and says so. `(x * 4) // 4` equals `x` in Python and does
39
+ **not** at 6 bits, because `x * 4` wraps. Both facts are worth knowing and
40
+ they are different facts.
41
+
42
+ Every operation whose meaning would be ambiguous is rejected rather than
43
+ guessed. `//` and `%` are supported only for constant powers of two, because
44
+ general division needs a restoring divider whose cost is rarely worth it and
45
+ whose silent wrong answer would be worse. Unsupported syntax raises
46
+ `UnsupportedConstruct` with the line number. A verifier that quietly ignores a
47
+ construct is worse than no verifier.
48
+ """
49
+
50
+ from __future__ import annotations
51
+
52
+ import ast
53
+ import inspect
54
+ import textwrap
55
+ from typing import Callable, Sequence
56
+
57
+ from dratify.cnf import CNF
58
+ from .encodings import Encoder
59
+ from dratify.lits import mk_lit, neg
60
+ from .native import available as _native_available
61
+ from dratify.proof import MemoryProof
62
+
63
+ __all__ = [
64
+ "equivalent",
65
+ "EquivalenceResult",
66
+ "UnsupportedConstruct",
67
+ "BitVec",
68
+ "compile_function",
69
+ ]
70
+
71
+
72
+ class UnsupportedConstruct(Exception):
73
+ """Raised for Python this module will not model.
74
+
75
+ Deliberately fatal. Silently approximating a construct would make every
76
+ "equivalent" answer meaningless, since you could never tell whether the
77
+ proof covered the code you wrote or a simplification of it.
78
+ """
79
+
80
+
81
+ # --------------------------------------------------------------------------
82
+ # bit vectors
83
+ # --------------------------------------------------------------------------
84
+
85
+
86
+ class BitVec:
87
+ """A fixed-width integer as a list of literals, least significant first.
88
+
89
+ Two's complement. Every operation wraps at `width` bits, which is what the
90
+ hardware a Python int eventually runs on does anyway -- the difference is
91
+ that here it is explicit and declared.
92
+ """
93
+
94
+ __slots__ = ("bits", "enc")
95
+
96
+ def __init__(self, enc: Encoder, bits: Sequence[int]) -> None:
97
+ self.enc = enc
98
+ self.bits = list(bits)
99
+
100
+ @property
101
+ def width(self) -> int:
102
+ return len(self.bits)
103
+
104
+ @classmethod
105
+ def constant(cls, enc: Encoder, value: int, width: int) -> "BitVec":
106
+ t, f = enc.true_lit, enc.false_lit
107
+ return cls(enc, [(t if (value >> i) & 1 else f) for i in range(width)])
108
+
109
+ @classmethod
110
+ def input(cls, enc: Encoder, width: int, name: str) -> "BitVec":
111
+ return cls(enc, [enc.new_lit(f"{name}[{i}]") for i in range(width)])
112
+
113
+ # -- bitwise --------------------------------------------------------
114
+
115
+ def _zip(self, other: "BitVec", op) -> "BitVec":
116
+ w = max(self.width, other.width)
117
+ a, b = self.extend(w), other.extend(w)
118
+ return BitVec(self.enc, [op(x, y) for x, y in zip(a.bits, b.bits)])
119
+
120
+ def __and__(self, o): return self._zip(o, lambda x, y: self.enc.and_gate([x, y]))
121
+ def __or__(self, o): return self._zip(o, lambda x, y: self.enc.or_gate([x, y]))
122
+ def __xor__(self, o): return self._zip(o, lambda x, y: self.enc.xor_gate(x, y))
123
+
124
+ def __invert__(self) -> "BitVec":
125
+ return BitVec(self.enc, [neg(b) for b in self.bits])
126
+
127
+ def extend(self, width: int) -> "BitVec":
128
+ """Sign-extend or truncate to `width`."""
129
+ if width == self.width:
130
+ return self
131
+ if width < self.width:
132
+ return BitVec(self.enc, self.bits[:width])
133
+ sign = self.bits[-1] if self.bits else self.enc.false_lit
134
+ return BitVec(self.enc, self.bits + [sign] * (width - self.width))
135
+
136
+ # -- arithmetic -----------------------------------------------------
137
+
138
+ def __add__(self, other: "BitVec") -> "BitVec":
139
+ """Ripple-carry addition, wrapping at the top bit."""
140
+ enc = self.enc
141
+ w = max(self.width, other.width)
142
+ a, b = self.extend(w), other.extend(w)
143
+ out, carry = [], enc.false_lit
144
+ for i in range(w):
145
+ x, y = a.bits[i], b.bits[i]
146
+ xy = enc.xor_gate(x, y)
147
+ out.append(enc.xor_gate(xy, carry))
148
+ carry = enc.or_gate([enc.and_gate([x, y]),
149
+ enc.and_gate([x, carry]),
150
+ enc.and_gate([y, carry])])
151
+ return BitVec(enc, out)
152
+
153
+ def __neg__(self) -> "BitVec":
154
+ """Two's complement negation: ~x + 1."""
155
+ return (~self) + BitVec.constant(self.enc, 1, self.width)
156
+
157
+ def __sub__(self, other: "BitVec") -> "BitVec":
158
+ return self + (-other.extend(max(self.width, other.width)))
159
+
160
+ def __mul__(self, other: "BitVec") -> "BitVec":
161
+ """Shift-and-add array multiplier, truncated to `width`.
162
+
163
+ Same construction as the factoring benchmark in `bench/run_bench.py`,
164
+ which is where its correctness was first exercised.
165
+ """
166
+ enc = self.enc
167
+ w = max(self.width, other.width)
168
+ a, b = self.extend(w), other.extend(w)
169
+ acc = BitVec.constant(enc, 0, w)
170
+ for i in range(w):
171
+ # partial product a << i, gated on b[i]
172
+ row = [enc.false_lit] * i + [
173
+ enc.and_gate([a.bits[j], b.bits[i]]) for j in range(w - i)
174
+ ]
175
+ acc = acc + BitVec(enc, row)
176
+ return acc.extend(w)
177
+
178
+ def shl(self, k: int) -> "BitVec":
179
+ if k >= self.width:
180
+ return BitVec.constant(self.enc, 0, self.width)
181
+ return BitVec(self.enc,
182
+ [self.enc.false_lit] * k + self.bits[: self.width - k])
183
+
184
+ def shr(self, k: int, arithmetic: bool = True) -> "BitVec":
185
+ """Right shift. Arithmetic (sign-propagating) by default, because
186
+ Python's `>>` on negative ints is arithmetic."""
187
+ fill = self.bits[-1] if arithmetic else self.enc.false_lit
188
+ if k >= self.width:
189
+ return BitVec(self.enc, [fill] * self.width)
190
+ return BitVec(self.enc, self.bits[k:] + [fill] * k)
191
+
192
+ # -- comparison -----------------------------------------------------
193
+
194
+ def eq(self, other: "BitVec") -> int:
195
+ w = max(self.width, other.width)
196
+ a, b = self.extend(w), other.extend(w)
197
+ same = [neg(self.enc.xor_gate(x, y)) for x, y in zip(a.bits, b.bits)]
198
+ return self.enc.and_gate(same)
199
+
200
+ def slt(self, other: "BitVec") -> int:
201
+ """Signed less-than, via the sign of the difference with overflow
202
+ correction: a < b iff (a-b) is negative, XOR the signed-overflow flag."""
203
+ enc = self.enc
204
+ w = max(self.width, other.width) + 1 # one extra bit kills the overflow case
205
+ a, b = self.extend(w), other.extend(w)
206
+ return (a - b).bits[-1]
207
+
208
+ def is_zero(self) -> int:
209
+ return self.enc.and_gate([neg(b) for b in self.bits])
210
+
211
+
212
+ def _ite_bv(enc: Encoder, cond: int, a: BitVec, b: BitVec) -> BitVec:
213
+ w = max(a.width, b.width)
214
+ a, b = a.extend(w), b.extend(w)
215
+ return BitVec(enc, [enc.ite(cond, x, y) for x, y in zip(a.bits, b.bits)])
216
+
217
+
218
+ # --------------------------------------------------------------------------
219
+ # the compiler
220
+ # --------------------------------------------------------------------------
221
+
222
+ _CMP = {ast.Eq, ast.NotEq, ast.Lt, ast.LtE, ast.Gt, ast.GtE}
223
+
224
+
225
+ class _Compiler(ast.NodeVisitor):
226
+ """Compile one function body into a circuit.
227
+
228
+ Control flow is handled by *path conditions* rather than by branching:
229
+ every assignment under an `if` becomes a select between the new value and
230
+ the old one, gated on the branch condition. That is what makes early
231
+ returns work -- each `return` records `(condition, value)`, and the result
232
+ is the chain of selects over those pairs.
233
+ """
234
+
235
+ def __init__(self, enc: Encoder, width: int, env: dict[str, BitVec]) -> None:
236
+ self.enc = enc
237
+ self.width = width
238
+ self.env = dict(env)
239
+ self.returns: list[tuple[int, BitVec]] = []
240
+ self.pc = enc.true_lit # current path condition
241
+
242
+ # -- helpers --------------------------------------------------------
243
+
244
+ def _fail(self, node: ast.AST, what: str) -> None:
245
+ line = getattr(node, "lineno", "?")
246
+ raise UnsupportedConstruct(
247
+ f"line {line}: {what} is not supported. This checker models a "
248
+ f"restricted subset on purpose -- approximating it silently would "
249
+ f"make every 'equivalent' answer meaningless."
250
+ )
251
+
252
+ def const(self, v: int) -> BitVec:
253
+ return BitVec.constant(self.enc, v, self.width)
254
+
255
+ # -- expressions ----------------------------------------------------
256
+
257
+ def expr(self, node: ast.AST) -> BitVec:
258
+ if isinstance(node, ast.Constant):
259
+ if isinstance(node.value, bool):
260
+ return self.const(1 if node.value else 0)
261
+ if isinstance(node.value, int):
262
+ return self.const(node.value)
263
+ self._fail(node, f"constant of type {type(node.value).__name__}")
264
+
265
+ if isinstance(node, ast.Name):
266
+ if node.id not in self.env:
267
+ self._fail(node, f"name {node.id!r} used before assignment")
268
+ return self.env[node.id]
269
+
270
+ if isinstance(node, ast.UnaryOp):
271
+ if isinstance(node.op, ast.USub):
272
+ return -self.expr(node.operand)
273
+ if isinstance(node.op, ast.UAdd):
274
+ return self.expr(node.operand)
275
+ if isinstance(node.op, ast.Invert):
276
+ return ~self.expr(node.operand)
277
+ if isinstance(node.op, ast.Not):
278
+ c = self.truth(node.operand)
279
+ return _ite_bv(self.enc, c, self.const(0), self.const(1))
280
+ self._fail(node, type(node.op).__name__)
281
+
282
+ if isinstance(node, ast.BinOp):
283
+ return self.binop(node)
284
+
285
+ if isinstance(node, (ast.Compare, ast.BoolOp)):
286
+ c = self.truth(node)
287
+ return _ite_bv(self.enc, c, self.const(1), self.const(0))
288
+
289
+ if isinstance(node, ast.IfExp):
290
+ c = self.truth(node.test)
291
+ return _ite_bv(self.enc, c, self.expr(node.body), self.expr(node.orelse))
292
+
293
+ if isinstance(node, ast.Call):
294
+ self._fail(node, "function calls (inline the callee, or pass it "
295
+ "through `helpers=`)")
296
+
297
+ self._fail(node, type(node).__name__)
298
+
299
+ def binop(self, node: ast.BinOp) -> BitVec:
300
+ op = node.op
301
+ left = self.expr(node.left)
302
+
303
+ # shifts and power-of-two division need a constant right operand
304
+ if isinstance(op, (ast.LShift, ast.RShift, ast.FloorDiv, ast.Mod)):
305
+ if not (isinstance(node.right, ast.Constant)
306
+ and isinstance(node.right.value, int)):
307
+ self._fail(node, f"{type(op).__name__} by a non-constant")
308
+ k = node.right.value
309
+ if isinstance(op, ast.LShift):
310
+ return left.shl(k)
311
+ if isinstance(op, ast.RShift):
312
+ return left.shr(k)
313
+ # // and % only by powers of two: a general divider is a large
314
+ # circuit and rarely what a refactor hinges on
315
+ if k <= 0 or (k & (k - 1)) != 0:
316
+ self._fail(node, f"{type(op).__name__} by {k} "
317
+ f"(only positive powers of two)")
318
+ shift = k.bit_length() - 1
319
+ if isinstance(op, ast.FloorDiv):
320
+ return left.shr(shift)
321
+ mask = self.const(k - 1)
322
+ return left & mask
323
+
324
+ right = self.expr(node.right)
325
+ if isinstance(op, ast.Add):
326
+ return left + right
327
+ if isinstance(op, ast.Sub):
328
+ return left - right
329
+ if isinstance(op, ast.Mult):
330
+ return left * right
331
+ if isinstance(op, ast.BitAnd):
332
+ return left & right
333
+ if isinstance(op, ast.BitOr):
334
+ return left | right
335
+ if isinstance(op, ast.BitXor):
336
+ return left ^ right
337
+ self._fail(node, type(op).__name__)
338
+
339
+ def truth(self, node: ast.AST) -> int:
340
+ """Compile an expression used as a condition into a single literal."""
341
+ if isinstance(node, ast.Compare):
342
+ if len(node.ops) != 1:
343
+ self._fail(node, "chained comparison")
344
+ op = node.ops[0]
345
+ if type(op) not in _CMP:
346
+ self._fail(node, type(op).__name__)
347
+ a, b = self.expr(node.left), self.expr(node.comparators[0])
348
+ if isinstance(op, ast.Eq):
349
+ return a.eq(b)
350
+ if isinstance(op, ast.NotEq):
351
+ return neg(a.eq(b))
352
+ if isinstance(op, ast.Lt):
353
+ return a.slt(b)
354
+ if isinstance(op, ast.GtE):
355
+ return neg(a.slt(b))
356
+ if isinstance(op, ast.Gt):
357
+ return b.slt(a)
358
+ if isinstance(op, ast.LtE):
359
+ return neg(b.slt(a))
360
+
361
+ if isinstance(node, ast.BoolOp):
362
+ parts = [self.truth(v) for v in node.values]
363
+ if isinstance(node.op, ast.And):
364
+ return self.enc.and_gate(parts)
365
+ return self.enc.or_gate(parts)
366
+
367
+ if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not):
368
+ return neg(self.truth(node.operand))
369
+
370
+ if isinstance(node, ast.Constant) and isinstance(node.value, bool):
371
+ return self.enc.true_lit if node.value else self.enc.false_lit
372
+
373
+ # any other expression is truthy when non-zero
374
+ return neg(self.expr(node).is_zero())
375
+
376
+ # -- statements -----------------------------------------------------
377
+
378
+ def block(self, body: list[ast.stmt]) -> None:
379
+ for stmt in body:
380
+ self.stmt(stmt)
381
+
382
+ def stmt(self, node: ast.stmt) -> None:
383
+ if isinstance(node, ast.Assign):
384
+ if len(node.targets) != 1 or not isinstance(node.targets[0], ast.Name):
385
+ self._fail(node, "tuple or attribute assignment")
386
+ self.assign(node.targets[0].id, self.expr(node.value))
387
+ return
388
+
389
+ if isinstance(node, ast.AugAssign):
390
+ if not isinstance(node.target, ast.Name):
391
+ self._fail(node, "augmented assignment to a non-name")
392
+ fake = ast.BinOp(left=ast.Name(id=node.target.id, ctx=ast.Load()),
393
+ op=node.op, right=node.value)
394
+ ast.copy_location(fake, node)
395
+ ast.copy_location(fake.left, node)
396
+ self.assign(node.target.id, self.binop(fake))
397
+ return
398
+
399
+ if isinstance(node, ast.Return):
400
+ if node.value is None:
401
+ # A bare `return` yields None. Modelling it as 0 makes a false
402
+ # proof against a function that really does return 0, and it
403
+ # contradicts `return None`, which this compiler rejects.
404
+ self._fail(node, "a bare `return` (it yields None, not an "
405
+ "integer; `return None` is rejected for the "
406
+ "same reason)")
407
+ value = self.expr(node.value)
408
+ self.returns.append((self.pc, value))
409
+ # everything after a return on this path is unreachable
410
+ self.pc = self.enc.false_lit
411
+ return
412
+
413
+ if isinstance(node, ast.If):
414
+ cond = self.truth(node.test)
415
+ before = dict(self.env)
416
+ outer_pc = self.pc
417
+
418
+ self.pc = self.enc.and_gate([outer_pc, cond])
419
+ self.block(node.body)
420
+ then_env, then_pc = dict(self.env), self.pc
421
+
422
+ self.env = dict(before)
423
+ self.pc = self.enc.and_gate([outer_pc, neg(cond)])
424
+ self.block(node.orelse)
425
+ else_env, else_pc = dict(self.env), self.pc
426
+
427
+ merged = {}
428
+ for name in set(then_env) | set(else_env):
429
+ t = then_env.get(name, before.get(name))
430
+ e = else_env.get(name, before.get(name))
431
+ if t is None or e is None:
432
+ continue # defined on only one branch: not readable after
433
+ merged[name] = t if t is e else _ite_bv(self.enc, cond, t, e)
434
+ self.env = merged
435
+ # the path survives if either branch did
436
+ self.pc = self.enc.or_gate([then_pc, else_pc])
437
+ return
438
+
439
+ if isinstance(node, ast.For):
440
+ self.unroll_for(node)
441
+ return
442
+
443
+ if isinstance(node, (ast.Pass, ast.Expr)) and not isinstance(
444
+ getattr(node, "value", None), ast.Call):
445
+ return # docstrings and bare expressions have no effect
446
+
447
+ if isinstance(node, ast.While):
448
+ self._fail(node, "while loops (bound is not statically known; use "
449
+ "`for i in range(k)` with a constant k)")
450
+
451
+ self._fail(node, type(node).__name__)
452
+
453
+ def assign(self, name: str, value: BitVec) -> None:
454
+ """Assign under the current path condition.
455
+
456
+ Inside a branch the assignment is conditional, so the variable becomes
457
+ a select between the new value and whatever it held before.
458
+ """
459
+ self.env[name] = value
460
+
461
+ def unroll_for(self, node: ast.For) -> None:
462
+ if not isinstance(node.target, ast.Name):
463
+ self._fail(node, "loop over a non-name target")
464
+ call = node.iter
465
+ if not (isinstance(call, ast.Call) and isinstance(call.func, ast.Name)
466
+ and call.func.id == "range"):
467
+ self._fail(node, "iteration over anything but range(...)")
468
+ args = []
469
+ for a in call.args:
470
+ if not (isinstance(a, ast.Constant) and isinstance(a.value, int)):
471
+ self._fail(node, "range() with a non-constant bound "
472
+ "(the loop must be statically bounded to unroll)")
473
+ args.append(a.value)
474
+ rng = range(*args) if args else range(0)
475
+ if len(rng) > 256:
476
+ self._fail(node, f"range of {len(rng)} iterations (limit 256; "
477
+ f"unrolling is linear in the bound)")
478
+ if node.orelse:
479
+ self._fail(node, "for/else")
480
+ for i in rng:
481
+ self.env[node.target.id] = self.const(i)
482
+ self.block(node.body)
483
+
484
+ # -- result ---------------------------------------------------------
485
+
486
+ def result(self, body: list[ast.stmt]) -> BitVec:
487
+ # This check must precede the empty-`returns` case below: a body with
488
+ # no return at all also falls off the end, and short-circuiting first
489
+ # would skip the guard entirely.
490
+ # The last return is only a sound unconditional fallback when control
491
+ # cannot reach the end of the body without it. If it can, the function
492
+ # yields None on that path and there is no bit-vector for that -- so
493
+ # refuse rather than silently model the guarded value as unconditional.
494
+ # Modelling it would make a *false proof*: two functions agreeing on
495
+ # the guarded value would be "proved equivalent" while differing
496
+ # wherever one falls through.
497
+ if not _definitely_returns(body):
498
+ raise UnsupportedConstruct(
499
+ "control can reach the end of the function without returning, "
500
+ "so it yields None on that path; pyeq models integers only. "
501
+ "Add an explicit return.")
502
+ if not self.returns: # unreachable once the guard above holds
503
+ return self.const(0)
504
+ value = self.returns[-1][1]
505
+ for cond, v in reversed(self.returns[:-1]):
506
+ value = _ite_bv(self.enc, cond, v, value)
507
+ return value
508
+
509
+
510
+ def _definitely_returns(body: list[ast.stmt]) -> bool:
511
+ """True when control cannot reach the end of `body` without returning.
512
+
513
+ A `for` loop never counts: its bound can be zero, so the body may not run.
514
+ An `if` counts only when both arms are present and both definitely return.
515
+ """
516
+ for stmt in body:
517
+ if isinstance(stmt, ast.Return):
518
+ return True
519
+ if isinstance(stmt, ast.If):
520
+ if (stmt.orelse and _definitely_returns(stmt.body)
521
+ and _definitely_returns(stmt.orelse)):
522
+ return True
523
+ return False
524
+
525
+
526
+ def _function_ast(fn: Callable) -> ast.FunctionDef:
527
+ try:
528
+ src = textwrap.dedent(inspect.getsource(fn))
529
+ except (OSError, TypeError) as e:
530
+ raise UnsupportedConstruct(
531
+ f"cannot read the source of {getattr(fn, '__name__', fn)!r}: {e}"
532
+ ) from None
533
+ tree = ast.parse(src)
534
+ for node in tree.body:
535
+ if isinstance(node, ast.FunctionDef):
536
+ return node
537
+ raise UnsupportedConstruct("no function definition found in the source")
538
+
539
+
540
+ def compile_function(
541
+ fn: Callable,
542
+ enc: Encoder,
543
+ widths: dict[str, int],
544
+ inputs: dict[str, BitVec] | None = None,
545
+ ) -> tuple[BitVec, dict[str, BitVec]]:
546
+ """Compile `fn` into a circuit. Returns (result, inputs used)."""
547
+ tree = _function_ast(fn)
548
+ if tree.args.vararg or tree.args.kwarg or tree.args.kwonlyargs:
549
+ raise UnsupportedConstruct("*args, **kwargs and keyword-only parameters")
550
+ if tree.args.defaults or tree.args.kw_defaults:
551
+ # Every parameter is compiled as a free symbolic input, so a default is
552
+ # simply dropped. Two functions with different defaults would then be
553
+ # "proved equivalent" while disagreeing on every call that omits the
554
+ # argument.
555
+ raise UnsupportedConstruct(
556
+ "default argument values are not modelled -- every parameter is "
557
+ "compiled as a free input, so the default is silently dropped")
558
+ names = [a.arg for a in tree.args.args]
559
+ missing = [n for n in names if n not in widths]
560
+ if missing:
561
+ raise UnsupportedConstruct(
562
+ f"no width declared for parameter(s) {', '.join(missing)}; "
563
+ f"pass widths={{'{missing[0]}': 8, ...}}"
564
+ )
565
+ width = max(widths[n] for n in names) if names else 8
566
+
567
+ env = dict(inputs) if inputs else {}
568
+ for n in names:
569
+ if n not in env:
570
+ env[n] = BitVec.input(enc, widths[n], n).extend(width)
571
+ c = _Compiler(enc, width, env)
572
+ c.block(tree.body)
573
+ return c.result(tree.body), {n: env[n] for n in names}
574
+
575
+
576
+ # --------------------------------------------------------------------------
577
+ # the miter
578
+ # --------------------------------------------------------------------------
579
+
580
+
581
+ class EquivalenceResult:
582
+ """Outcome of an equivalence check.
583
+
584
+ ``proved`` is three-valued and the distinction matters:
585
+
586
+ ``True`` the functions agree on every input at the declared widths, and
587
+ (unless ``verify=False``) a DRAT proof of that was replayed by
588
+ an independent checker.
589
+ ``False`` they differ, and ``counterexample`` is an input where they do.
590
+ ``None`` the conflict budget ran out. Nothing was decided either way.
591
+
592
+ ``None`` is falsy, so ``if result:`` is still safe -- an undecided check
593
+ never reads as a proof. Code that needs the distinction must test
594
+ ``is True`` / ``is None`` explicitly.
595
+ """
596
+
597
+ __slots__ = ("proved", "counterexample", "width", "vars", "clauses",
598
+ "seconds", "outputs", "python_outputs", "overflow_only",
599
+ "proof_checked", "proof_steps", "conflicts")
600
+
601
+ def __init__(self) -> None:
602
+ self.proved: bool | None = False
603
+ #: True when a DRAT proof of the UNSAT miter was replayed and accepted
604
+ self.proof_checked: bool = False
605
+ #: length of that proof, in steps
606
+ self.proof_steps: int = 0
607
+ self.conflicts: int = 0
608
+ self.counterexample: dict[str, int] | None = None
609
+ #: what the two circuits produce at the declared width
610
+ self.outputs: tuple[int, int] | None = None
611
+ #: what the two Python functions produce with arbitrary precision
612
+ self.python_outputs: tuple[int, int] | None = None
613
+ #: True when the circuits differ but Python agrees -- the divergence is
614
+ #: an artefact of fixed-width wrapping, not of the refactor
615
+ self.overflow_only: bool = False
616
+ self.width: int = 0
617
+ self.vars: int = 0
618
+ self.clauses: int = 0
619
+ self.seconds: float = 0.0
620
+
621
+ def __bool__(self) -> bool:
622
+ # `is True`, not truthiness: an exhausted budget must never read as a
623
+ # proof, and `None` would be the value most likely to be mistaken for
624
+ # one by code that only ever checks `if result:`
625
+ return self.proved is True
626
+
627
+ def report(self) -> str:
628
+ if self.proved is None:
629
+ return (f"c UNDECIDED at {self.width} bits: budget exhausted after "
630
+ f"{self.conflicts} conflicts. Nothing was proved and "
631
+ f"nothing was refuted; raise max_conflicts or narrow the "
632
+ f"widths.")
633
+ if self.proved:
634
+ how = (f"proof of {self.proof_steps} steps verified"
635
+ if self.proof_checked else "UNVERIFIED: verify=False")
636
+ return (f"c equivalent at {self.width} bits, {how} "
637
+ f"({self.vars} vars, {self.clauses} clauses, "
638
+ f"{self.seconds*1000:.0f} ms)")
639
+ args = ", ".join(f"{k}={v}" for k, v in (self.counterexample or {}).items())
640
+ got = ("" if self.outputs is None
641
+ else f" -> {self.outputs[0]} vs {self.outputs[1]}")
642
+ head = f"c NOT equivalent at {self.width} bits: {args}{got}"
643
+ if self.overflow_only:
644
+ head += (f"\nc ...but Python agrees here "
645
+ f"({self.python_outputs[0]}): the difference is fixed-width "
646
+ f"overflow, not the refactor. Widen, or accept it as a real "
647
+ f"difference for machine integers.")
648
+ return head
649
+
650
+
651
+ def _as_signed(bits: list[bool]) -> int:
652
+ """Interpret a little-endian bit list as two's complement."""
653
+ n = 0
654
+ for i, b in enumerate(bits):
655
+ if b:
656
+ n |= 1 << i
657
+ # fall through
658
+ if bits and bits[-1]:
659
+ n -= 1 << len(bits)
660
+ return n
661
+
662
+
663
+ class ProofRejected(RuntimeError):
664
+ """The solver said UNSAT and the independent checker disagreed.
665
+
666
+ This is never a statement about the user's code. It means the solver, the
667
+ checker or the circuit compiler is broken, and it is raised rather than
668
+ returned because there is no honest verdict to hand back: the two things
669
+ that are supposed to agree did not.
670
+ """
671
+
672
+
673
+ def _solve_miter(formula: CNF, engine: str, budget: int | None, want_proof: bool):
674
+ """Solve the miter, optionally logging a DRAT proof.
675
+
676
+ Returns ``(status, model, steps, conflicts)`` with ``status is None``
677
+ meaning the budget ran out.
678
+
679
+ This deliberately does not go through :func:`cdclkit.pipeline.solve_adaptive`.
680
+ Preprocessing rewrites the formula -- bounded variable elimination in
681
+ particular adds clauses that do not follow from the input alone -- so a
682
+ refutation of the preprocessed formula is not a refutation of the miter,
683
+ and checking it against the miter would fail. Given a choice between a
684
+ faster solve and a checkable one, this module takes the checkable one.
685
+ """
686
+ if engine == "native" and _native_available():
687
+ from . import native
688
+ from dratify.lits import from_dimacs
689
+
690
+ s = native.require().Solver(formula.nvars)
691
+ if want_proof:
692
+ s.enable_proof() # must precede the first clause
693
+ for c in formula.clauses:
694
+ if not s.add_clause(list(c)):
695
+ return False, None, None, s.conflicts
696
+ res = s.solve(budget)
697
+ if res is None:
698
+ return None, None, None, s.conflicts
699
+ steps = (None if res or not want_proof else
700
+ [(k, tuple(from_dimacs(d) for d in lits))
701
+ for k, lits in s.proof_steps()])
702
+ return res, (list(s.model) if res else None), steps, s.conflicts
703
+
704
+ from .solver import Solver
705
+
706
+ proof = MemoryProof() if want_proof else None
707
+ s = Solver(formula.nvars, proof=proof)
708
+ if not s.add_cnf(formula):
709
+ return False, None, (proof.steps if proof else None), s.stats.conflicts
710
+ res = s.solve(max_conflicts=budget)
711
+ if res is None:
712
+ return None, None, None, s.stats.conflicts
713
+ steps = None if res or not want_proof else proof.steps
714
+ return res, (list(s.model) if res else None), steps, s.stats.conflicts
715
+
716
+
717
+ def equivalent(
718
+ f: Callable,
719
+ g: Callable,
720
+ widths: dict[str, int],
721
+ engine: str = "native",
722
+ verify: bool = True,
723
+ max_conflicts: int | None = None,
724
+ ) -> EquivalenceResult:
725
+ """Prove `f` and `g` agree on every input, or return one where they differ.
726
+
727
+ `widths` maps each parameter name to its bit width. Both functions must
728
+ take the same parameters.
729
+
730
+ `verify` (on by default) makes the solver emit a DRAT proof of the
731
+ equivalence and replays it through an independent checker before
732
+ `result.proved` is allowed to be True. It costs roughly the solve again.
733
+ Turning it off means the answer rests on the solver's word, which is the
734
+ one thing this project exists not to ask of anyone.
735
+
736
+ `max_conflicts` bounds the search. On exhaustion the result is
737
+ `proved=None` -- undecided -- never `True`.
738
+
739
+ The answer is about **fixed-width two's-complement arithmetic** at the
740
+ declared widths, not about Python's arbitrary-precision integers -- see the
741
+ module docstring. A counterexample is re-simulated before being returned,
742
+ so it is never spurious.
743
+ """
744
+ import time
745
+
746
+ t0 = time.perf_counter()
747
+ formula = CNF()
748
+ enc = Encoder(formula)
749
+
750
+ out_f, inputs = compile_function(f, enc, widths)
751
+ out_g, _ = compile_function(g, enc, widths, inputs=inputs)
752
+
753
+ w = max(out_f.width, out_g.width)
754
+ a, b = out_f.extend(w), out_g.extend(w)
755
+ # the miter: assert at least one output bit differs
756
+ enc.add([enc.xor_gate(x, y) for x, y in zip(a.bits, b.bits)])
757
+
758
+ r = EquivalenceResult()
759
+ r.width = max(widths.values()) if widths else 0
760
+ r.vars = formula.nvars
761
+ r.clauses = formula.nclauses
762
+
763
+ status, model, steps, conflicts = _solve_miter(
764
+ formula, engine, max_conflicts, want_proof=verify)
765
+ r.conflicts = conflicts
766
+
767
+ if status is None: # budget exhausted -- decided nothing
768
+ r.proved = None
769
+ r.seconds = time.perf_counter() - t0
770
+ return r
771
+
772
+ if status is False:
773
+ # The miter is unsatisfiable: no input distinguishes the two
774
+ # functions. That is the claim the whole call exists to make, so it is
775
+ # the claim that gets checked rather than trusted.
776
+ if verify:
777
+ from dratify.proof import check_proof
778
+
779
+ chk = check_proof(formula, steps or [])
780
+ if not chk.ok:
781
+ raise ProofRejected(
782
+ f"the solver reported the functions equivalent, and the "
783
+ f"independent checker rejected its proof at step "
784
+ f"{chk.failed_step}: {chk.reason}. This is a bug in cdclkit, "
785
+ f"not in your code. Please report it with both functions "
786
+ f"and the widths."
787
+ )
788
+ r.proof_checked = True
789
+ r.proof_steps = chk.steps
790
+ r.proved = True
791
+ r.seconds = time.perf_counter() - t0
792
+ return r
793
+
794
+ r.seconds = time.perf_counter() - t0
795
+
796
+ def value(bv: BitVec, n: int) -> int:
797
+ return _as_signed([model[l >> 1] != bool(l & 1) for l in bv.bits[:n]])
798
+
799
+ r.counterexample = {
800
+ name: value(bv, widths[name]) for name, bv in inputs.items()
801
+ }
802
+ # What the circuits actually produce. This is the authoritative answer:
803
+ # it is the semantics the proof is about.
804
+ r.outputs = (value(a, w), value(b, w))
805
+ if r.outputs[0] == r.outputs[1]:
806
+ raise AssertionError(
807
+ f"the solver reported a difference at {r.counterexample} but both "
808
+ f"circuits evaluate to {r.outputs[0]} there. That is a bug in this "
809
+ f"compiler, not in your code."
810
+ )
811
+
812
+ # And what Python says, which is *not* the same question: Python integers
813
+ # are arbitrary precision, so an operation that overflows the declared width
814
+ # wraps in the circuit and does not in Python. When the circuits differ but
815
+ # Python agrees, the divergence is an overflow artefact -- still a real
816
+ # difference for fixed-width machine integers, but a different finding, and
817
+ # the caller should be told which one they have.
818
+ try:
819
+ fv, gv = f(**r.counterexample), g(**r.counterexample)
820
+ r.python_outputs = (fv, gv)
821
+ r.overflow_only = (fv == gv)
822
+ except Exception:
823
+ r.python_outputs = None # not every function accepts the raw ints
824
+ return r