igneum/tools/attack/f1-shadow/z3check.py
2026-10-07 10:02:44 +00:00

177 lines
6 KiB
Python

#!/usr/bin/env python3
"""Attack-pass F1: z3 equivalence of the harness's normal-form DAG against the straight-line shadow block.
Reads the JSON `attack-f1 windows` writes (a list of windows: the instruction slice, the reachable DAG nodes and
the 8 output node ids) and proves, per window, that for every input register file and every `sel` the DAG's
outputs equal the instruction-by-instruction run (the verifier's `step` semantics, verify.rs). A window with a
`shfl` is modelled on all 32 lanes, any other on one lane (every other op is lane-local).
z3check.py FILE.json [--timeout-ms 60000] [--lanes 32]
One line per window: seed, start, instrs, lanes, nodes, result (proved | COUNTEREXAMPLE | unknown), seconds.
Exit 1 on any COUNTEREXAMPLE.
"""
import json
import sys
import time
import z3
W = 32
def rotl(x, n):
n %= 32
return z3.RotateLeft(x, n) if n else x
def rotr_var(x, s):
# the kernels' rotr_var: amount & 31; z3's variable rotate matches rotate_right for 0..31
return z3.RotateRight(x, s & 31)
def mulhi(a, b):
p = z3.ZeroExt(32, a) * z3.ZeroExt(32, b)
return z3.Extract(63, 32, p)
def straight(instrs, regs, sel, lanes):
r = [list(v) for v in regs]
for ins in instrs:
d, a, b = ins["dst"], ins["src"], ins["src2"]
op = ins["op"]
if op == "add":
for l in range(lanes):
c = z3.If(z3.Extract(ins["bit"], ins["bit"], sel[l]) == 1, z3.BitVecVal(ins["imm2"], W), z3.BitVecVal(ins["imm"], W))
r[d][l] = r[d][l] + r[a][l] + c
elif op == "sub":
for l in range(lanes):
r[d][l] = r[d][l] - r[a][l]
elif op == "mul":
for l in range(lanes):
r[d][l] = r[d][l] * r[a][l]
elif op == "mulhi":
for l in range(lanes):
r[d][l] = mulhi(r[d][l], r[a][l])
elif op == "xor":
for l in range(lanes):
r[d][l] = r[d][l] ^ r[a][l]
elif op == "or":
for l in range(lanes):
r[d][l] = r[d][l] | r[a][l]
elif op == "rotl":
for l in range(lanes):
r[d][l] = rotl(r[d][l], ins["rot"])
elif op == "rotr":
for l in range(lanes):
r[d][l] = rotr_var(r[d][l], r[a][l])
elif op == "mad":
for l in range(lanes):
r[d][l] = r[a][l] * r[b][l] + r[d][l]
elif op == "shfl":
src = list(r[a])
m = ins["mask"]
for l in range(lanes):
r[d][l] = r[d][l] ^ src[l ^ m]
else:
raise SystemExit("not an ALU op: " + op)
return r
def dag_eval(nodes, regs, sel, lanes):
byid = {n["id"]: n for n in nodes}
memo = {}
def get(i):
if i in memo:
return memo[i]
n = byid[i]
k = n["k"]
if k == "in":
v = list(regs[n["r"]])
elif k == "csel":
v = [z3.If(z3.Extract(n["bit"], n["bit"], sel[l]) == 1, z3.BitVecVal(n["imm2"], W), z3.BitVecVal(n["imm"], W)) for l in range(lanes)]
elif k == "zero":
v = [z3.BitVecVal(0, W) for _ in range(lanes)]
elif k == "sum":
v = [z3.BitVecVal(0, W) for _ in range(lanes)]
for t, c in n["t"]:
tv = get(t)
for l in range(lanes):
v[l] = v[l] + tv[l] * z3.BitVecVal(c, W)
elif k == "xor":
v = [z3.BitVecVal(0, W) for _ in range(lanes)]
for t, rot, m in n["t"]:
tv = get(t)
for l in range(lanes):
v[l] = v[l] ^ rotl(tv[(l ^ m) % lanes], rot)
elif k == "or":
v = [z3.BitVecVal(0, W) for _ in range(lanes)]
for t in n["t"]:
tv = get(t)
for l in range(lanes):
v[l] = v[l] | tv[l]
elif k == "rotr":
xv, sv = get(n["x"]), get(n["s"])
v = [rotr_var(xv[l], (sv[l] & 31) * n["n"]) for l in range(lanes)]
elif k == "mul":
av, bv = get(n["a"]), get(n["b"])
v = [z3.ZeroExt(32, av[l]) * z3.ZeroExt(32, bv[l]) for l in range(lanes)]
elif k == "lo":
pv = get(n["p"])
v = [z3.Extract(31, 0, pv[l]) for l in range(lanes)]
elif k == "hi":
pv = get(n["p"])
v = [z3.Extract(63, 32, pv[l]) for l in range(lanes)]
else:
raise SystemExit("unknown node kind " + k)
v = [z3.simplify(x) for x in v]
memo[i] = v
return v
return get
def main():
args = sys.argv[1:]
path = args[0]
timeout = 60000
force_lanes = None
if "--timeout-ms" in args:
timeout = int(args[args.index("--timeout-ms") + 1])
if "--lanes" in args:
force_lanes = int(args[args.index("--lanes") + 1])
windows = json.load(open(path))
bad = 0
for w in windows:
has_shfl = any(i["op"] == "shfl" for i in w["instrs"])
lanes = force_lanes or (32 if has_shfl else 1)
regs = [[z3.BitVec(f"r{r}_{l}", W) for l in range(lanes)] for r in range(8)]
sel = [z3.BitVec(f"sel_{l}", W) for l in range(lanes)]
t0 = time.time()
sl = straight(w["instrs"], regs, sel, lanes)
get = dag_eval(w["nodes"], regs, sel, lanes)
s = z3.Solver()
s.set("timeout", timeout)
diffs = []
for r, oid in enumerate(w["outputs"]):
dv = get(oid)
for l in range(lanes):
diffs.append(dv[l] != sl[r][l])
s.add(z3.Or(diffs))
res = s.check()
dt = time.time() - t0
if res == z3.unsat:
verdict = "proved"
elif res == z3.sat:
verdict = "COUNTEREXAMPLE"
bad += 1
else:
verdict = "unknown"
print(f"{w['seed']} start={w['start']} instrs={len(w['instrs'])} lanes={lanes} nodes={len(w['nodes'])} {verdict} {dt:.2f}s", flush=True)
print(f"windows {len(windows)} counterexamples {bad}")
sys.exit(1 if bad else 0)
if __name__ == "__main__":
main()