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/__init__.py +152 -0
- cdclkit/__main__.py +10 -0
- cdclkit/brute.py +210 -0
- cdclkit/cli.py +513 -0
- cdclkit/encodings.py +842 -0
- cdclkit/heap.py +180 -0
- cdclkit/model.py +420 -0
- cdclkit/mus.py +159 -0
- cdclkit/native.py +111 -0
- cdclkit/pipeline.py +212 -0
- cdclkit/portfolio.py +683 -0
- cdclkit/preprocess.py +500 -0
- cdclkit/pyeq.py +824 -0
- cdclkit/solver.py +1377 -0
- cdclkit-0.1.0.dist-info/METADATA +136 -0
- cdclkit-0.1.0.dist-info/RECORD +20 -0
- cdclkit-0.1.0.dist-info/WHEEL +5 -0
- cdclkit-0.1.0.dist-info/entry_points.txt +2 -0
- cdclkit-0.1.0.dist-info/licenses/LICENSE +202 -0
- cdclkit-0.1.0.dist-info/top_level.txt +1 -0
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
|