177 lines
6 KiB
Python
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()
|