#!/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()