1244 lines
53 KiB
Python
1244 lines
53 KiB
Python
#!/usr/bin/env python3
|
|
"""attack-f2: SAT differential and linear trail search on 1 to 4 keyed applications of the Igneum mixer
|
|
(igneum-pow/src/memhard.rs `mixer`), attack pass row F2 of docs/plans/cryptanalysis.md 4.2.
|
|
|
|
One application on 16 words of 32 bits: per word (s ^ (RC + rk)) * MUL (odd), then a ChaCha-shaped double round
|
|
with the four column rotations ROT[0..3] and the four diagonal rotations ROT[4..7]. The op list below is the ONE
|
|
description of that sequence; the value evaluator (checked against `attack-f2 vectors`, the Rust ground truth)
|
|
and both CNF builders consume it, so what the SAT models is what the code does.
|
|
|
|
Models (bit-exact where stated, trail models otherwise):
|
|
differential (XOR differences): XOR with a constant: free. Odd multiply: the XOR difference is converted to a
|
|
modular difference (exact: each set bit below the MSB is a sign choice at probability 1/2), multiplied by MUL
|
|
(exact, a circuit on the difference variables), and converted back to an XOR difference (exact: a carry chain,
|
|
one bit of weight per position where the difference bit and the carry differ). The two conversions are
|
|
treated as independent (the usual Markov assumption). Modular addition: Lipmaa-Moriai (exact). XOR, rotation:
|
|
linear. The MSB passes the multiply and every addition for free; the rotations are what move it.
|
|
linear (masks): modular addition by the exact carry-mask automaton (per bit: a carry mask bit sigma; sigma = 1
|
|
costs one bit of correlation and needs an odd mask weight on (carry, x, y); sigma = 0 needs all three zero;
|
|
bit 0 and a known-zero carry take any mask at cost 1), checked against brute force in `selftest`. The odd
|
|
multiply is modelled as its shift-and-add decomposition (the set bits of MUL), each partial sum an addition
|
|
under the same automaton with the known-zero low bits of a shifted copy transparent. Trails only: the
|
|
decomposition's partial sums are correlated, so a found trail is verified on the real code and the bound is a
|
|
trail bound under this decomposition.
|
|
rotational-XOR is measured on the real code by `attack-f2 rx` / `rx-word`, not modelled here.
|
|
|
|
Usage:
|
|
model.py selftest the evaluator against the Rust vectors, the add models
|
|
against brute force (n = 8), the mult model against sampling
|
|
model.py search --kind diff|lin --day D --variant real|rot0|nomul --apps K --state FILE --budget SEC
|
|
[--cap W] [--solver cadical195] resumable bound search; writes FILE (JSON) and FILE.trail
|
|
model.py show --state FILE the state in one line
|
|
The parameter files come from `attack-f2 params` (ROT, MUL); vectors from `attack-f2 vectors`.
|
|
"""
|
|
import argparse
|
|
import itertools
|
|
import math
|
|
import json
|
|
import os
|
|
import random
|
|
import sys
|
|
import time
|
|
|
|
MASK32 = 0xFFFFFFFF
|
|
N = 32
|
|
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
# Parameters and the op list
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
|
|
|
|
def load_params(path):
|
|
d = {}
|
|
for line in open(path):
|
|
parts = line.split()
|
|
if not parts:
|
|
continue
|
|
d[parts[0]] = parts[1:]
|
|
rot = [int(x) for x in d["rot"]]
|
|
mul = [int(x, 16) for x in d["mul"]]
|
|
rc = [int(x, 16) for x in d["rc"]]
|
|
rk = [int(x, 16) for x in d["rk_round0"]]
|
|
return {"day": d["day"][0], "variant": d["variant"][0], "rot": rot, "mul": mul, "rc": rc, "rk": rk}
|
|
|
|
|
|
QR_COLS = [(0, 4, 8, 12), (1, 5, 9, 13), (2, 6, 10, 14), (3, 7, 11, 15)]
|
|
QR_DIAGS = [(0, 5, 10, 15), (1, 6, 11, 12), (2, 7, 8, 13), (3, 4, 9, 14)]
|
|
|
|
|
|
def ops_one_application(p, rk):
|
|
"""The op list of one application: ('xorc', i, C), ('mul', i, MUL), then 8 quarter rounds of
|
|
('add', a, b), ('xor', d, a), ('rot', d, r1), ('add', c, d), ('xor', b, c), ('rot', b, r2), (twice)."""
|
|
ops = []
|
|
for i in range(16):
|
|
ops.append(("xorc", i, (p["rc"][i] + rk) & MASK32))
|
|
ops.append(("mul", i, p["mul"][i]))
|
|
r = p["rot"]
|
|
|
|
def qr(a, b, c, d, r1, r2, r3, r4):
|
|
ops.extend([("add", a, b), ("xor", d, a), ("rot", d, r1), ("add", c, d), ("xor", b, c), ("rot", b, r2)])
|
|
ops.extend([("add", a, b), ("xor", d, a), ("rot", d, r3), ("add", c, d), ("xor", b, c), ("rot", b, r4)])
|
|
|
|
for (a, b, c, d) in QR_COLS:
|
|
qr(a, b, c, d, r[0], r[1], r[2], r[3])
|
|
for (a, b, c, d) in QR_DIAGS:
|
|
qr(a, b, c, d, r[4], r[5], r[6], r[7])
|
|
return ops
|
|
|
|
|
|
def rotl(x, r):
|
|
r %= 32
|
|
return ((x << r) | (x >> (32 - r))) & MASK32 if r else x
|
|
|
|
|
|
def evaluate(state, ops):
|
|
s = list(state)
|
|
for op in ops:
|
|
k = op[0]
|
|
if k == "xorc":
|
|
s[op[1]] ^= op[2]
|
|
elif k == "mul":
|
|
s[op[1]] = (s[op[1]] * op[2]) & MASK32
|
|
elif k == "add":
|
|
s[op[1]] = (s[op[1]] + s[op[2]]) & MASK32
|
|
elif k == "xor":
|
|
s[op[1]] ^= s[op[2]]
|
|
elif k == "rot":
|
|
s[op[1]] = rotl(s[op[1]], op[2])
|
|
else:
|
|
raise ValueError(k)
|
|
return s
|
|
|
|
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
# CNF builder
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
|
|
|
|
class CNF:
|
|
def __init__(self):
|
|
self.nv = 1
|
|
self.clauses = []
|
|
self.TRUE = 1
|
|
self.clauses.append([1])
|
|
|
|
@property
|
|
def FALSE(self):
|
|
return -self.TRUE
|
|
|
|
def var(self):
|
|
self.nv += 1
|
|
return self.nv
|
|
|
|
def vars(self, n):
|
|
return [self.var() for _ in range(n)]
|
|
|
|
def add(self, *lits):
|
|
self.clauses.append(list(lits))
|
|
|
|
def xor2(self, a, b):
|
|
"""z = a xor b."""
|
|
z = self.var()
|
|
self.add(-z, a, b)
|
|
self.add(-z, -a, -b)
|
|
self.add(z, -a, b)
|
|
self.add(z, a, -b)
|
|
return z
|
|
|
|
def parity(self, lits, value, guard=None):
|
|
"""Forbid every assignment of `lits` whose XOR differs from `value` (when `guard` is true, if given)."""
|
|
n = len(lits)
|
|
for bits in itertools.product([0, 1], repeat=n):
|
|
if (sum(bits) & 1) != value:
|
|
cl = [(-l if b else l) for l, b in zip(lits, bits)]
|
|
if guard is not None:
|
|
cl.append(-guard)
|
|
self.add(*cl)
|
|
|
|
def xor3(self, a, b, c):
|
|
z = self.var()
|
|
self.parity([a, b, c, z], 0)
|
|
return z
|
|
|
|
def maj(self, a, b, c):
|
|
z = self.var()
|
|
self.add(-z, a, b)
|
|
self.add(-z, a, c)
|
|
self.add(-z, b, c)
|
|
self.add(z, -a, -b)
|
|
self.add(z, -a, -c)
|
|
self.add(z, -b, -c)
|
|
return z
|
|
|
|
def xor_vec(self, xs, ys):
|
|
return [self.xor2(a, b) for a, b in zip(xs, ys)]
|
|
|
|
def xor_many(self, vecs):
|
|
if not vecs:
|
|
return [self.FALSE] * N
|
|
acc = vecs[0]
|
|
for v in vecs[1:]:
|
|
acc = self.xor_vec(acc, v)
|
|
return acc
|
|
|
|
def or_clause(self, lits):
|
|
self.add(*lits)
|
|
|
|
def ripple_add(self, xs, ys, cin):
|
|
"""z = x + y + cin mod 2^32, bit-exact on variables (no probability)."""
|
|
z = []
|
|
c = cin
|
|
for i in range(N):
|
|
z.append(self.xor3(xs[i], ys[i], c))
|
|
if i < N - 1:
|
|
c = self.maj(xs[i], ys[i], c)
|
|
return z
|
|
|
|
def const_vec(self, v):
|
|
return [self.TRUE if (v >> i) & 1 else self.FALSE for i in range(N)]
|
|
|
|
|
|
def naf(c):
|
|
"""Non-adjacent form digits of c as a list of (position, +1/-1), low to high, for c odd and < 2^32.
|
|
The product c * x mod 2^32 only needs digits below 32."""
|
|
digits = []
|
|
k = 0
|
|
while c:
|
|
if c & 1:
|
|
d = 2 - (c & 3) # 1 if c mod 4 == 1 else -1
|
|
digits.append((k, d))
|
|
c -= d
|
|
c >>= 1
|
|
k += 1
|
|
return [(k, d) for (k, d) in digits if k < 32]
|
|
|
|
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
# Differential model
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
|
|
|
|
class DiffModel:
|
|
def __init__(self, cnf, family="general"):
|
|
self.cnf = cnf
|
|
self.family = family # "general": the Markov multiply model; "msb": every word difference at a multiply is 0 or MSB (exact)
|
|
self.weights = [] # weight literals (each true literal costs one bit of probability)
|
|
self.weight_tags = []
|
|
self.weight_app = [] # the application index of each weight literal (Matsui split)
|
|
self.mul_io = [] # (application index, word, input difference lits, output difference lits)
|
|
self.app_index = 0
|
|
|
|
def add_lm(self, xs, ys, tag="add"):
|
|
"""Lipmaa-Moriai: XOR differences xs, ys -> zs through modular addition."""
|
|
c = self.cnf
|
|
zs = c.vars(N)
|
|
# bit 0: z0 = x0 ^ y0
|
|
z0 = c.xor2(xs[0], ys[0])
|
|
c.add(-zs[0], z0)
|
|
c.add(zs[0], -z0)
|
|
for i in range(N - 1):
|
|
w = c.var()
|
|
# not w -> x_i = y_i = z_i
|
|
c.add(w, -xs[i], ys[i])
|
|
c.add(w, xs[i], -ys[i])
|
|
c.add(w, -ys[i], zs[i])
|
|
c.add(w, ys[i], -zs[i])
|
|
# w -> not all equal
|
|
c.add(-w, xs[i], ys[i], zs[i])
|
|
c.add(-w, -xs[i], -ys[i], -zs[i])
|
|
# not w -> x_{i+1} ^ y_{i+1} ^ z_{i+1} ^ x_i = 0
|
|
c.parity([xs[i + 1], ys[i + 1], zs[i + 1], xs[i]], 0, guard=-w)
|
|
self.weights.append(w)
|
|
self.weight_tags.append(tag)
|
|
self.weight_app.append(self.app_index)
|
|
return zs
|
|
|
|
def mul(self, ds, mulc, tag="mul"):
|
|
"""XOR difference ds through x -> (x ^ K) * mulc: XOR -> modular (sign choices), times mulc, modular -> XOR."""
|
|
c = self.cnf
|
|
if self.family == "msb":
|
|
# exact: (x ^ 2^31) * c = x * c ^ 2^31 for odd c, and no other nonzero difference passes with probability 1
|
|
for i in range(N - 1):
|
|
c.add(-ds[i])
|
|
return list(ds)
|
|
# sign split: pos_i xor neg_i = d_i, not both (i < 31); bit 31: pos = d_31
|
|
pos, neg = [], []
|
|
for i in range(N - 1):
|
|
p, q = c.var(), c.var()
|
|
c.add(-ds[i], p, q)
|
|
c.add(-p, ds[i])
|
|
c.add(-q, ds[i])
|
|
c.add(-p, -q)
|
|
pos.append(p)
|
|
neg.append(q)
|
|
self.weights.append(ds[i])
|
|
self.weight_tags.append(tag + "_in")
|
|
self.weight_app.append(self.app_index)
|
|
pos.append(ds[N - 1])
|
|
neg.append(c.FALSE)
|
|
# delta = pos - neg = pos + ~neg + 1
|
|
delta = c.ripple_add(pos, [-q for q in neg], c.TRUE)
|
|
# delta' = mulc * delta by NAF shift-add-subtract (exact arithmetic on variables)
|
|
digits = naf(mulc)
|
|
assert digits[0][0] == 0
|
|
acc = None
|
|
for (k, d) in digits:
|
|
term = [c.FALSE] * k + delta[: N - k]
|
|
if acc is None:
|
|
if d == 1:
|
|
acc = term
|
|
else:
|
|
acc = c.ripple_add([-t for t in term], c.const_vec(0), c.TRUE) # -term
|
|
elif d == 1:
|
|
acc = c.ripple_add(acc, term, c.FALSE)
|
|
else:
|
|
acc = c.ripple_add(acc, [-t for t in term], c.TRUE)
|
|
dp = acc
|
|
# modular dp on y -> XOR difference: carries k_0 = 0; w_i = (dp_i != k_i); not w_i -> k_{i+1} = dp_i
|
|
out = []
|
|
kprev = c.FALSE
|
|
for i in range(N):
|
|
out.append(c.xor2(dp[i], kprev))
|
|
if i < N - 1:
|
|
w = c.xor2(dp[i], kprev)
|
|
knext = c.var()
|
|
c.add(w, -knext, dp[i])
|
|
c.add(w, knext, -dp[i])
|
|
self.weights.append(w)
|
|
self.weight_tags.append(tag + "_out")
|
|
self.weight_app.append(self.app_index)
|
|
kprev = knext
|
|
return out
|
|
|
|
def application(self, state, ops):
|
|
s = list(state)
|
|
for op in ops:
|
|
k = op[0]
|
|
if k == "xorc":
|
|
pass
|
|
elif k == "mul":
|
|
din = s[op[1]]
|
|
s[op[1]] = self.mul(din, op[2], tag="mul%d" % op[1])
|
|
self.mul_io.append((self.app_index, op[1], din, s[op[1]]))
|
|
elif k == "add":
|
|
s[op[1]] = self.add_lm(s[op[1]], s[op[2]])
|
|
elif k == "xor":
|
|
s[op[1]] = self.cnf.xor_vec(s[op[1]], s[op[2]])
|
|
elif k == "rot":
|
|
r = op[2] % 32
|
|
v = s[op[1]]
|
|
s[op[1]] = v[N - r:] + v[: N - r] if r else v
|
|
return s
|
|
|
|
|
|
def build_diff(p, apps, family="general"):
|
|
cnf = CNF()
|
|
m = DiffModel(cnf, family=family)
|
|
s0 = [cnf.vars(N) for _ in range(16)]
|
|
cnf.or_clause([b for w in s0 for b in w]) # nonzero input difference
|
|
bounds = [s0]
|
|
s = s0
|
|
for j in range(apps):
|
|
m.app_index = j + 1
|
|
s = m.application(s, ops_one_application(p, p["rk"][j]))
|
|
bounds.append(s)
|
|
cnf.or_clause([b for w in s for b in w]) # nonzero output difference
|
|
return cnf, m, bounds
|
|
|
|
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
# Linear model (masks), built backwards over an SSA graph of the op list
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
|
|
|
|
class LinModel:
|
|
def __init__(self, cnf, family="general"):
|
|
self.cnf = cnf
|
|
self.family = family # "general": shift-and-add trail model; "low2": output masks of a multiply within bits 0 and 1 (exact)
|
|
self.weights = []
|
|
self.weight_tags = []
|
|
self.mul_io = [] # (application index, word, input mask lits, output mask lits)
|
|
self.add_io = [] # (u lits, v lits, w lits, k0) of every modular addition, for the per-adder signed hull check
|
|
self.weight_app = []
|
|
self.weight_idx = []
|
|
self.cur_app = 0
|
|
self.cur_idx = 0
|
|
self.qr_masks = []
|
|
|
|
def add_automaton(self, w, k0=0, tag="add"):
|
|
"""Masks (u on x, v on y) for output mask w through z = x + y with the low k0 bits of y known zero
|
|
(z_i = x_i for i < k0). Returns (u, v). Weight = number of carry-mask bits sigma_{i} set, i in k0+1..31."""
|
|
c = self.cnf
|
|
u, v = c.vars(N), c.vars(N)
|
|
self.add_io.append((u, v, w, k0, self.cur_idx, self.cur_app))
|
|
# transparent low bits: u_i = w_i, v_i free, no carry
|
|
for i in range(k0):
|
|
c.add(-u[i], w[i])
|
|
c.add(u[i], -w[i])
|
|
sigma = {k0: c.FALSE} # carry mask into bit k0 is on a constant-zero carry
|
|
for i in range(k0, N):
|
|
s_in = sigma[i]
|
|
if i == N - 1:
|
|
s_out = c.FALSE
|
|
else:
|
|
s_out = c.var()
|
|
sigma[i + 1] = s_out
|
|
self.weights.append(s_out)
|
|
self.weight_tags.append(tag)
|
|
self.weight_app.append(self.cur_app)
|
|
self.weight_idx.append(self.cur_idx)
|
|
a = c.xor2(u[i], w[i])
|
|
b = c.xor2(v[i], w[i])
|
|
if i == k0:
|
|
# carry into this bit is constant 0: p is irrelevant; s_out = 0 -> a = b = 0; s_out = 1 -> free
|
|
c.add(s_out, -a)
|
|
c.add(s_out, -b)
|
|
else:
|
|
pbit = c.xor2(s_in, w[i])
|
|
# s_out = 0 -> p = a = b = 0
|
|
c.add(s_out, -pbit)
|
|
c.add(s_out, -a)
|
|
c.add(s_out, -b)
|
|
# s_out = 1 -> p ^ a ^ b = 1
|
|
c.parity([pbit, a, b], 1, guard=s_out)
|
|
return u, v
|
|
|
|
def build(self, p, apps):
|
|
"""Returns (input masks, output masks, boundary masks list) as literal vectors. Every use of a value carries
|
|
the index of the op that consumes it; the mask ON a value (for the constraint at the op producing it) is the
|
|
XOR of all its uses, while the mask of a value AT a boundary (an application boundary, a quarter-round
|
|
boundary) is the XOR of its uses at or after that boundary (a value read inside its own application, such
|
|
as the final `a` of a quarter round read by `d ^= a`, keeps that use out of the boundary mask)."""
|
|
c = self.cnf
|
|
nxt = [16]
|
|
slots = list(range(16))
|
|
ssa = []
|
|
boundaries = [(0, list(slots))] # (op index of the boundary, slot values at it)
|
|
qrs = [] # (app, qr index, start op index, slots before, end op index, slots after)
|
|
|
|
def new():
|
|
v = nxt[0]
|
|
nxt[0] += 1
|
|
return v
|
|
|
|
for j in range(apps):
|
|
ops = ops_one_application(p, p["rk"][j])
|
|
qr_start = None
|
|
n_in_qr = 0
|
|
for op in ops:
|
|
k = op[0]
|
|
idx = len(ssa)
|
|
if k in ("add",) and n_in_qr == 0:
|
|
qr_start = (idx, list(slots))
|
|
if k == "xorc":
|
|
o = new()
|
|
ssa.append(("xorc", o, slots[op[1]], op[2], j + 1, op[1]))
|
|
slots[op[1]] = o
|
|
elif k == "mul":
|
|
o = new()
|
|
ssa.append(("mul", o, slots[op[1]], op[2], j + 1, op[1]))
|
|
slots[op[1]] = o
|
|
elif k == "add":
|
|
o = new()
|
|
ssa.append(("add", o, slots[op[1]], slots[op[2]], j + 1))
|
|
slots[op[1]] = o
|
|
elif k == "xor":
|
|
o = new()
|
|
ssa.append(("xor", o, slots[op[1]], slots[op[2]], j + 1))
|
|
slots[op[1]] = o
|
|
elif k == "rot":
|
|
o = new()
|
|
ssa.append(("rot", o, slots[op[1]], op[2] % 32, j + 1))
|
|
slots[op[1]] = o
|
|
if k in ("add", "xor", "rot"):
|
|
n_in_qr += 1
|
|
if n_in_qr == 12:
|
|
qrs.append((j + 1, len(qrs) % 8, qr_start[0], qr_start[1], len(ssa), list(slots)))
|
|
n_in_qr = 0
|
|
boundaries.append((len(ssa), list(slots)))
|
|
nvals = nxt[0]
|
|
INF = 10 ** 9
|
|
uses = [[] for _ in range(nvals)] # (consumer op index, mask vector)
|
|
out_masks = []
|
|
for v in slots:
|
|
m = c.vars(N)
|
|
uses[v].append((INF, m))
|
|
out_masks.append(m)
|
|
value_mask = [None] * nvals
|
|
|
|
def mask_of(v):
|
|
if value_mask[v] is None:
|
|
value_mask[v] = c.xor_many([m for (_, m) in uses[v]]) if uses[v] else [c.FALSE] * N
|
|
return value_mask[v]
|
|
|
|
for idx in range(len(ssa) - 1, -1, -1):
|
|
op = ssa[idx]
|
|
k, o = op[0], op[1]
|
|
self.cur_app = op[4]
|
|
self.cur_idx = idx
|
|
w = mask_of(o)
|
|
if k == "xorc":
|
|
uses[op[2]].append((idx, w))
|
|
elif k == "rot":
|
|
r = op[3]
|
|
uses[op[2]].append((idx, [w[(m + r) % N] for m in range(N)]))
|
|
elif k == "xor":
|
|
uses[op[2]].append((idx, w))
|
|
uses[op[3]].append((idx, w))
|
|
elif k == "add":
|
|
u, v = self.add_automaton(w)
|
|
uses[op[2]].append((idx, u))
|
|
uses[op[3]].append((idx, v))
|
|
elif k == "mul" and self.family == "low2":
|
|
mulc = op[3]
|
|
for i in range(2, N):
|
|
c.add(-w[i])
|
|
u = [c.FALSE] * N
|
|
u[1] = w[1]
|
|
u[0] = c.xor2(w[0], w[1]) if (mulc >> 1) & 1 else w[0]
|
|
uses[op[2]].append((idx, u))
|
|
self.mul_io.append((op[4], op[5], u, w))
|
|
elif k == "mul":
|
|
mulc = op[3]
|
|
bits = [i for i in range(N) if (mulc >> i) & 1]
|
|
assert bits[0] == 0
|
|
cur = w
|
|
contributions = []
|
|
for kk in reversed(bits[1:]):
|
|
u, v = self.add_automaton(cur, k0=kk, tag="mul%d" % 0)
|
|
contributions.append([v[m + kk] for m in range(N - kk)] + [c.FALSE] * kk)
|
|
cur = u
|
|
contributions.append(cur)
|
|
for cv in contributions:
|
|
uses[op[2]].append((idx, cv))
|
|
self.mul_io.append((op[4], op[5], mask_of(op[2]), w))
|
|
|
|
def mask_at(v, at):
|
|
ms = [m for (i, m) in uses[v] if i >= at]
|
|
return c.xor_many(ms) if ms else [c.FALSE] * N
|
|
|
|
in_masks = [mask_of(v) for v in range(16)]
|
|
c.or_clause([b for m in in_masks for b in m])
|
|
c.or_clause([b for m in out_masks for b in m])
|
|
bmasks = [[mask_at(v, at) for v in b] for (at, b) in boundaries]
|
|
self.qr_masks = [(app, q, [mask_at(v, s0) for v in sb], [mask_at(v, e0) for v in se]) for (app, q, s0, sb, e0, se) in qrs]
|
|
return in_masks, out_masks, bmasks
|
|
|
|
|
|
def build_lin(p, apps, family="general"):
|
|
cnf = CNF()
|
|
m = LinModel(cnf, family=family)
|
|
in_masks, out_masks, bmasks = m.build(p, apps)
|
|
return cnf, m, bmasks
|
|
|
|
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
# Search
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
|
|
|
|
def solve_with_bound(cnf, weights, bound, solver_name, deadline, conf_chunk=200000):
|
|
"""SAT (model), UNSAT (None), or 'timeout'."""
|
|
from pysat.card import CardEnc, EncType
|
|
from pysat.solvers import Solver
|
|
|
|
top = cnf.nv
|
|
enc = EncType.seqcounter if len(weights) * max(bound, 1) <= 4_000_000 else EncType.totalizer
|
|
card = CardEnc.atmost(lits=weights, bound=bound, top_id=top, encoding=enc)
|
|
with Solver(name=solver_name, bootstrap_with=cnf.clauses) as s:
|
|
for cl in card.clauses:
|
|
s.add_clause(cl)
|
|
while True:
|
|
try:
|
|
s.conf_budget(conf_chunk)
|
|
r = s.solve_limited()
|
|
except NotImplementedError:
|
|
r = s.solve()
|
|
if r is True:
|
|
return s.get_model()
|
|
if r is False:
|
|
return None
|
|
if time.time() > deadline:
|
|
return "timeout"
|
|
|
|
|
|
def qr_ops(p, app, q):
|
|
"""The 12 ops of quarter round q (0..7) of one application, as an op list on the 16-word state."""
|
|
ops = ops_one_application(p, p["rk"][app - 1])
|
|
arx = [op for op in ops if op[0] in ("add", "xor", "rot")]
|
|
return arx[12 * q: 12 * q + 12]
|
|
|
|
|
|
def qr_model_weight(m, model_set, app, q):
|
|
"""The sigma bits set in the four adders of quarter round q (0..7) of application app."""
|
|
arx_adds = sorted(i for (_u, _v, _w, k0, i, a) in m.add_io if k0 == 0 and a == app)
|
|
group = set(arx_adds[4 * q: 4 * q + 4])
|
|
return sum(1 for lit, i in zip(m.weights, m.weight_idx) if i in group and lit in model_set)
|
|
|
|
|
|
def qr_sample_corr(p, app, q, mi, mo, n):
|
|
import random as _r
|
|
rnd = _r.Random(0xF2 + app * 8 + q)
|
|
ops = qr_ops(p, app, q)
|
|
tot = 0
|
|
for _ in range(n):
|
|
x = [rnd.getrandbits(32) for _ in range(16)]
|
|
y = evaluate(x, ops)
|
|
par = 0
|
|
for i in range(16):
|
|
if mi[i]:
|
|
par ^= bin(mi[i] & x[i]).count("1") & 1
|
|
if mo[i]:
|
|
par ^= bin(mo[i] & y[i]).count("1") & 1
|
|
tot += -1 if par else 1
|
|
return tot / n
|
|
|
|
|
|
def verify_trail(binary, kind, trail_words, day, variant, log2, p=None):
|
|
"""Run `attack-f2 verify-diff|verify-lin` on the real code. Returns (chain_weight, [per-app weights]).
|
|
`binary == "python"`: the same measurement with this file's evaluator at 2^14 samples (local tests only)."""
|
|
import subprocess
|
|
import tempfile
|
|
if binary == "python":
|
|
import random as _r
|
|
rnd = _r.Random(11)
|
|
n = 1 << 14
|
|
k = len(trail_words) - 1
|
|
apps_ops = [ops_one_application(p, p["rk"][j]) for j in range(k)]
|
|
per = []
|
|
for j in range(k):
|
|
tot = 0
|
|
for _ in range(n):
|
|
x = [rnd.getrandbits(32) for _ in range(16)]
|
|
if kind == "diff":
|
|
y0 = evaluate(x, apps_ops[j])
|
|
y1 = evaluate([x[i] ^ trail_words[j][i] for i in range(16)], apps_ops[j])
|
|
tot += all((y0[i] ^ y1[i]) == trail_words[j + 1][i] for i in range(16))
|
|
else:
|
|
y = evaluate(x, apps_ops[j])
|
|
par = 0
|
|
for i in range(16):
|
|
par ^= bin(trail_words[j][i] & x[i]).count("1") & 1
|
|
par ^= bin(trail_words[j + 1][i] & y[i]).count("1") & 1
|
|
tot += -1 if par else 1
|
|
per.append(-math.log2(abs(tot) / n) if tot else float("inf"))
|
|
tot = 0
|
|
allops = [op for ops in apps_ops for op in ops]
|
|
for _ in range(n):
|
|
x = [rnd.getrandbits(32) for _ in range(16)]
|
|
if kind == "diff":
|
|
y0 = evaluate(x, allops)
|
|
y1 = evaluate([x[i] ^ trail_words[0][i] for i in range(16)], allops)
|
|
tot += all((y0[i] ^ y1[i]) == trail_words[k][i] for i in range(16))
|
|
else:
|
|
y = evaluate(x, allops)
|
|
par = 0
|
|
for i in range(16):
|
|
par ^= bin(trail_words[0][i] & x[i]).count("1") & 1
|
|
par ^= bin(trail_words[k][i] & y[i]).count("1") & 1
|
|
tot += -1 if par else 1
|
|
chain = -math.log2(abs(tot) / n) if tot else float("inf")
|
|
return chain, per
|
|
with tempfile.NamedTemporaryFile("w", suffix=".trail", delete=False) as f:
|
|
for t in trail_words:
|
|
f.write(" ".join(f"{x:08x}" for x in t) + "\n")
|
|
path = f.name
|
|
cmd = [binary, "verify-" + kind, "--day", day, "--variant", variant, "--trail", path, "--log2", str(log2)]
|
|
out = subprocess.run(cmd, capture_output=True, text=True, check=True).stdout
|
|
os.unlink(path)
|
|
per, chain = [], None
|
|
for line in out.splitlines():
|
|
f = line.split()
|
|
if not f:
|
|
continue
|
|
if f[0] == "app":
|
|
per.append(float(f[f.index("weight") + 1]) if kind == "diff" else float(f[f.index("abs_log2") + 1]))
|
|
elif f[0] == "chain":
|
|
chain = float(f[f.index("weight") + 1]) if kind == "diff" else float(f[f.index("abs_log2") + 1])
|
|
return chain, per
|
|
|
|
|
|
def lits_to_words(model_set, vecs):
|
|
out = []
|
|
for v in vecs:
|
|
x = 0
|
|
for i, lit in enumerate(v):
|
|
val = (lit in model_set) if lit > 0 else ((-lit) not in model_set)
|
|
if val:
|
|
x |= 1 << i
|
|
out.append(x)
|
|
return out
|
|
|
|
|
|
def next_probe(st, cap):
|
|
"""The next weight bound to ask, or None when the search is closed."""
|
|
lo = st["unsat_upto"] + 1 # smallest weight not yet excluded
|
|
if st.get("stuck_at") is not None:
|
|
return None
|
|
if st["sat_at"] is None:
|
|
ladder = [0, 1, 2, 3, 4, 6, 8, 12, 16, 20, 24, 28, 32, 40, 48, 56, 64, 80, 96, 112, 128, 160, 192, 256, 384, 512]
|
|
for w in ladder:
|
|
if w >= lo and w <= cap:
|
|
return w
|
|
return None if lo > cap else cap
|
|
if st["sat_at"] == lo:
|
|
return None
|
|
return (lo + st["sat_at"]) // 2
|
|
|
|
|
|
def cmd_search(args):
|
|
p = load_params(args.params)
|
|
st = {
|
|
"kind": args.kind, "family": args.family, "day": p["day"], "variant": p["variant"], "apps": args.apps, "rot": p["rot"],
|
|
"unsat_upto": -1, "sat_at": None, "stuck_at": None, "trail": None, "weight_by_tag": None,
|
|
"solver_seconds": 0.0, "log": [],
|
|
}
|
|
if os.path.exists(args.state):
|
|
st = json.load(open(args.state))
|
|
t0 = time.time()
|
|
deadline = t0 + args.budget
|
|
build_t = time.time()
|
|
if args.kind == "diff":
|
|
cnf, m, bounds = build_diff(p, args.apps, family=args.family)
|
|
else:
|
|
cnf, m, bounds = build_lin(p, args.apps, family=args.family)
|
|
if args.kind == "lin":
|
|
for (app, q, hi, ho) in st.get("blocked_qr", []):
|
|
for (a2, q2, min_, mout_) in m.qr_masks:
|
|
if (a2, q2) != (app, q):
|
|
continue
|
|
cl = []
|
|
for vec, hv in zip(min_ + mout_, hi + ho):
|
|
val = int(hv, 16)
|
|
for i, lit in enumerate(vec):
|
|
cl.append(-lit if (val >> i) & 1 else lit)
|
|
cnf.add(*cl)
|
|
for (jj, hin, hout) in st.get("blocked_app", []):
|
|
cl = []
|
|
for vec, hv in zip(bounds[jj - 1] + bounds[jj], hin + hout):
|
|
val = int(hv, 16)
|
|
for i, lit in enumerate(vec):
|
|
cl.append(-lit if (val >> i) & 1 else lit)
|
|
cnf.add(*cl)
|
|
for (bin_, bout) in st.get("blocked_boundary", []):
|
|
cl = []
|
|
for vec, hv in zip(bounds[0] + bounds[-1], bin_ + bout):
|
|
val = int(hv, 16)
|
|
for i, lit in enumerate(vec):
|
|
cl.append(-lit if (val >> i) & 1 else lit)
|
|
cnf.add(*cl)
|
|
if args.kind == "lin":
|
|
for (hu, hv, hw) in st.get("blocked", []):
|
|
uu, vv, wv = int(hu, 16), int(hv, 16), int(hw, 16)
|
|
for (u, v, ww, k0, _i, _a) in m.add_io:
|
|
cl = []
|
|
for vec, val in ((u, uu), (v, vv), (ww, wv)):
|
|
for i, lit in enumerate(vec):
|
|
cl.append(-lit if (val >> i) & 1 else lit)
|
|
cnf.add(*cl)
|
|
if args.per_app_min > 0:
|
|
from pysat.card import CardEnc as _CE2
|
|
for j in range(1, args.apps + 1):
|
|
lits = [l for l, a in zip(m.weights, m.weight_app) if a == j]
|
|
assert lits, "no weights tagged for application %d" % j
|
|
card = _CE2.atleast(lits=lits, bound=args.per_app_min, top_id=cnf.nv)
|
|
cnf.nv = max(cnf.nv, card.nv)
|
|
for cl in card.clauses:
|
|
cnf.clauses.append(cl)
|
|
st["per_app_min"] = args.per_app_min
|
|
st["unsat_upto"] = max(st["unsat_upto"], args.per_app_min * args.apps - 1)
|
|
build_t = time.time() - build_t
|
|
print(f"built {args.kind}/{args.family} apps={args.apps} vars={cnf.nv} clauses={len(cnf.clauses)} weights={len(m.weights)} in {build_t:.1f}s", flush=True)
|
|
while True:
|
|
w = next_probe(st, args.cap)
|
|
if w is None:
|
|
st["done"] = True
|
|
break
|
|
if time.time() > deadline - 5:
|
|
break
|
|
t1 = time.time()
|
|
r = solve_with_bound(cnf, m.weights, w, args.solver, deadline)
|
|
# linear trails: a trail whose correlation cancels inside one adder's hull (two carry-mask paths of
|
|
# opposite sign) is not a trail; block that adder's mask triple and ask again at the same bound
|
|
while args.kind == "lin" and r not in (None, "timeout"):
|
|
ms = set(l for l in r if l > 0)
|
|
bad = None
|
|
hull_w = 0.0
|
|
for (u, v, ww, k0, _i, _a) in m.add_io:
|
|
uu, vv, wv = lits_to_words(ms, [u, v, ww])
|
|
if uu == 0 and vv == 0 and wv == 0:
|
|
continue
|
|
h = dp_lin_corr(N, uu, vv, wv, k0=k0)
|
|
if h == 0:
|
|
bad = (u, v, ww, uu, vv, wv)
|
|
break
|
|
hull_w += -math.log2(abs(h))
|
|
if bad is None:
|
|
st["hull_weight"] = round(hull_w, 3)
|
|
break
|
|
u, v, ww, uu, vv, wv = bad
|
|
cl = []
|
|
for vec, val in ((u, uu), (v, vv), (ww, wv)):
|
|
for i, lit in enumerate(vec):
|
|
cl.append(-lit if (val >> i) & 1 else lit)
|
|
cnf.add(*cl)
|
|
st.setdefault("blocked", []).append([f"{uu:08x}", f"{vv:08x}", f"{wv:08x}"])
|
|
st["cancelled"] = st.get("cancelled", 0) + 1
|
|
print(f"bound {w}: cancelled trail (adder hull 0 at u {uu:08x} v {vv:08x} w {wv:08x}), blocked; {st['cancelled']} so far", flush=True)
|
|
save_state(st, args.state)
|
|
r = solve_with_bound(cnf, m.weights, w, args.solver, deadline)
|
|
# linear: every quarter round's sub-trail of small model weight is measured by sampling (the adders of one
|
|
# quarter round are dependent, so the piling-up product can be wrong); a sub-trail that does not hold is
|
|
# blocked at the quarter-round level, which removes every chain through it
|
|
while args.kind == "lin" and r not in (None, "timeout"):
|
|
ms = set(l for l in r if l > 0)
|
|
bad = None
|
|
for (app, q, min_, mout_) in m.qr_masks:
|
|
mi, mo = lits_to_words(ms, min_), lits_to_words(ms, mout_)
|
|
if not any(mi) and not any(mo):
|
|
continue
|
|
wq = qr_model_weight(m, ms, app, q)
|
|
if wq > 4:
|
|
continue
|
|
meas = qr_sample_corr(p, app, q, mi, mo, 1 << 16)
|
|
# resolvable: a correlation of 2^-(wq + 2) reads 4 sigma at 2^16 samples when wq <= 4
|
|
if abs(meas) < 2.0 ** -(wq + 2):
|
|
bad = (app, q, min_, mout_, mi, mo, wq, meas)
|
|
break
|
|
if bad is None:
|
|
break
|
|
app, q, min_, mout_, mi, mo, wq, meas = bad
|
|
cl = []
|
|
for vec, val in zip(min_ + mout_, mi + mo):
|
|
for i, lit in enumerate(vec):
|
|
cl.append(-lit if (val >> i) & 1 else lit)
|
|
cnf.add(*cl)
|
|
st.setdefault("blocked_qr", []).append([app, q, [f"{x:08x}" for x in mi], [f"{x:08x}" for x in mo]])
|
|
st["refuted_qr"] = st.get("refuted_qr", 0) + 1
|
|
print(f"bound {w}: quarter round app {app} qr {q} model weight {wq} measured |corr| 2^-{-math.log2(abs(meas)) if meas else 99:.1f}: blocked ({st['refuted_qr']} so far)", flush=True)
|
|
save_state(st, args.state)
|
|
r = solve_with_bound(cnf, m.weights, w, args.solver, deadline)
|
|
# every found trail is measured on the real code, application by application (and the chain): an
|
|
# application whose measured weight falls more than 2 bits short of its model weight does not hold, its
|
|
# boundary pair (difference or mask in, out) is blocked, and the solver is asked again at the same bound.
|
|
# Applications above --verify-max bits are not measurable by sampling and stay model-only.
|
|
while args.verifier and r not in (None, "timeout"):
|
|
ms = set(l for l in r if l > 0)
|
|
per_model = [sum(1 for lit, a in zip(m.weights, m.weight_app) if a == j and lit in ms) for j in range(1, args.apps + 1)]
|
|
total = sum(per_model)
|
|
checkable = [j for j in range(args.apps) if per_model[j] <= args.verify_max]
|
|
if not checkable:
|
|
break
|
|
trail = [lits_to_words(ms, b) for b in bounds]
|
|
wmax = max(per_model[j] for j in checkable)
|
|
if total <= args.verify_max:
|
|
wmax = max(wmax, total)
|
|
log2 = max(min(wmax + 8, 28) if args.kind == "diff" else min(2 * (wmax + 2) + 6, 28), 20)
|
|
chain, per = verify_trail(args.verifier, args.kind, trail, p["day"], p["variant"], log2, p=p)
|
|
bad = None
|
|
for j in checkable:
|
|
meas = per[j]
|
|
events = (2.0 ** log2) * (2.0 ** -meas) if args.kind == "diff" else (2.0 ** (log2 / 2.0)) * (2.0 ** -meas)
|
|
if meas > per_model[j] + 2.0 or events < (20.0 if args.kind == "diff" else 5.0):
|
|
bad = j
|
|
break
|
|
chain_ok = None
|
|
if total <= args.verify_max and chain is not None:
|
|
ev = (2.0 ** log2) * (2.0 ** -chain) if args.kind == "diff" else (2.0 ** (log2 / 2.0)) * (2.0 ** -chain)
|
|
chain_ok = chain <= total + 2.0 and ev >= (20.0 if args.kind == "diff" else 5.0)
|
|
st.setdefault("verifications", []).append({"bound": w, "model_total": total, "model_per_app": per_model, "log2": log2, "chain": chain, "per_app": per, "bad_app": None if bad is None else bad + 1, "chain_ok": chain_ok})
|
|
if bad is None and chain_ok is not False:
|
|
st["verified_chain_weight"] = chain if chain_ok else None
|
|
st["verified_per_app"] = per
|
|
st["verified_log2"] = log2
|
|
break
|
|
j = bad if bad is not None else 0
|
|
# block application j's boundary pair (all its applications' pairs when only the chain failed)
|
|
blocks = [j] if bad is not None else list(range(args.apps))
|
|
for jj in blocks:
|
|
cl = []
|
|
for vec, val in zip(bounds[jj] + bounds[jj + 1], trail[jj] + trail[jj + 1]):
|
|
for i, lit in enumerate(vec):
|
|
cl.append(-lit if (val >> i) & 1 else lit)
|
|
cnf.add(*cl)
|
|
st.setdefault("blocked_app", []).append([jj + 1, [f"{x:08x}" for x in trail[jj]], [f"{x:08x}" for x in trail[jj + 1]]])
|
|
st["refuted"] = st.get("refuted", 0) + 1
|
|
why = f"application {j + 1} model {per_model[j]} measured {per[j]}" if bad is not None else f"chain model {total} measured {chain}"
|
|
print(f"bound {w}: trail does not hold ({why} over 2^{log2}): blocked ({st['refuted']} so far)", flush=True)
|
|
save_state(st, args.state)
|
|
r = solve_with_bound(cnf, m.weights, w, args.solver, deadline)
|
|
dt = time.time() - t1
|
|
st["solver_seconds"] += dt
|
|
if r == "timeout":
|
|
st["log"].append({"bound": w, "result": "timeout", "secs": round(dt, 1), "utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())})
|
|
print(f"bound {w}: timeout after {dt:.0f}s (budget chunk)", flush=True)
|
|
# a chunk ran out: resume at the same bound next time unless the total cap says stuck
|
|
st["pending"] = w
|
|
break
|
|
if r is None:
|
|
st["unsat_upto"] = max(st["unsat_upto"], w)
|
|
st["log"].append({"bound": w, "result": "unsat", "secs": round(dt, 1)})
|
|
print(f"bound {w}: UNSAT in {dt:.1f}s", flush=True)
|
|
else:
|
|
ms = set(l for l in r if l > 0)
|
|
trail = [lits_to_words(ms, b) for b in bounds]
|
|
# weight by tag
|
|
by = {}
|
|
for lit, tag in zip(m.weights, m.weight_tags):
|
|
if lit in ms or (lit < 0 and -lit not in ms):
|
|
t = tag.split("_")[0] if tag.startswith("mul") and args.kind == "diff" else tag
|
|
t = "mul" if t.startswith("mul") else t
|
|
by[t] = by.get(t, 0) + 1
|
|
total = sum(by.values())
|
|
st["sat_at"] = w if st["sat_at"] is None else min(st["sat_at"], w)
|
|
if st["trail"] is None or total <= st.get("trail_weight", 10**9):
|
|
st["trail"] = [[f"{x:08x}" for x in t] for t in trail]
|
|
st["trail_weight"] = total
|
|
st["weight_by_tag"] = by
|
|
if st["trail_weight"] == total:
|
|
st["mults"] = [(j, i, f"{lits_to_words(ms, [a])[0]:08x}", f"{lits_to_words(ms, [b])[0]:08x}") for (j, i, a, b) in m.mul_io
|
|
if lits_to_words(ms, [a])[0] or lits_to_words(ms, [b])[0]]
|
|
st["log"].append({"bound": w, "result": "sat", "weight": total, "secs": round(dt, 1)})
|
|
print(f"bound {w}: SAT weight {total} {by} in {dt:.1f}s", flush=True)
|
|
# the found weight itself is an upper bound
|
|
st["sat_at"] = min(st["sat_at"], total)
|
|
st["pending"] = None
|
|
save_state(st, args.state)
|
|
st["elapsed_total"] = st.get("elapsed_total", 0.0) + (time.time() - t0)
|
|
save_state(st, args.state)
|
|
if st["trail"]:
|
|
with open(args.state + ".trail", "w") as f:
|
|
f.write(f"# {args.kind}/{args.family} day {p['day']} variant {p['variant']} apps {args.apps} weight {st['trail_weight']} by {st['weight_by_tag']} verified_chain {st.get('verified_chain_weight')} per_app {st.get('verified_per_app')}\n")
|
|
for t in st["trail"]:
|
|
f.write(" ".join(t) + "\n")
|
|
with open(args.state + ".mults", "w") as f:
|
|
f.write(f"# {args.kind}/{args.family} multiply-layer word transitions of the trail: app word in out\n")
|
|
for (j, i, a, b) in st.get("mults", []):
|
|
f.write(f"app {j} word {i} {a} {b}\n")
|
|
print(show_line(st), flush=True)
|
|
|
|
|
|
def save_state(st, path):
|
|
"""Atomic: a reader never sees a half-written state."""
|
|
tmp = path + ".tmp"
|
|
with open(tmp, "w") as f:
|
|
json.dump(st, f, indent=1)
|
|
os.replace(tmp, path)
|
|
|
|
|
|
def show_line(st):
|
|
best = st["trail_weight"] if st.get("trail") else None
|
|
return (f"{st['kind']}/{st.get('family', 'general')} day {st['day']} variant {st['variant']} apps {st['apps']}: best_found_weight={best} "
|
|
f"no_trail_at_or_below={st['unsat_upto']} done={st.get('done', False)} pending={st.get('pending')} "
|
|
f"per_app_min={st.get('per_app_min', 0)} hull_weight={st.get('hull_weight')} cancelled={st.get('cancelled', 0)} refuted={st.get('refuted', 0)} verified_chain={st.get('verified_chain_weight')} solver_s={st['solver_seconds']:.0f} by_tag={st.get('weight_by_tag')}")
|
|
|
|
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
# Self tests
|
|
# ----------------------------------------------------------------------------------------------------------------
|
|
|
|
|
|
def brute_lin_corr(n, u, v, w):
|
|
tot = 0
|
|
for x in range(1 << n):
|
|
for y in range(1 << n):
|
|
z = (x + y) & ((1 << n) - 1)
|
|
par = (bin(u & x).count("1") + bin(v & y).count("1") + bin(w & z).count("1")) & 1
|
|
tot += -1 if par else 1
|
|
return tot / (1 << (2 * n))
|
|
|
|
|
|
def dp_lin_corr(n, u, v, w, k0=0):
|
|
"""The automaton's signed trail sum (the exact hull), the rule the SAT encodes."""
|
|
from fractions import Fraction
|
|
# transparent low bits
|
|
for i in range(k0):
|
|
if ((u >> i) & 1) != ((w >> i) & 1):
|
|
return 0.0
|
|
states = {0: Fraction(1)} # sigma_in -> amplitude
|
|
for i in range(k0, n):
|
|
a = ((u >> i) & 1) ^ ((w >> i) & 1)
|
|
b = ((v >> i) & 1) ^ ((w >> i) & 1)
|
|
wi = (w >> i) & 1
|
|
nxt = {}
|
|
for s_in, amp in states.items():
|
|
p = s_in ^ wi
|
|
for s_out in ([0] if i == n - 1 else [0, 1]):
|
|
if i == k0:
|
|
if s_out == 0:
|
|
if a == 0 and b == 0:
|
|
nxt[0] = nxt.get(0, 0) + amp
|
|
else:
|
|
# AND correlations: masks (a,b): 00:+1/2 10:+1/2 01:+1/2 11:-1/2
|
|
sign = -1 if (a and b) else 1
|
|
nxt[1] = nxt.get(1, 0) + amp * Fraction(sign, 2)
|
|
else:
|
|
wt = p + a + b
|
|
if s_out == 0:
|
|
if wt == 0:
|
|
nxt[0] = nxt.get(0, 0) + amp
|
|
else:
|
|
if wt == 1:
|
|
nxt[1] = nxt.get(1, 0) + amp * Fraction(1, 2)
|
|
elif wt == 3:
|
|
nxt[1] = nxt.get(1, 0) + amp * Fraction(-1, 2)
|
|
states = nxt
|
|
return float(states.get(0, 0))
|
|
|
|
|
|
def brute_diff_prob(n, a, b, g):
|
|
cnt = 0
|
|
for x in range(1 << n):
|
|
for y in range(1 << n):
|
|
if (((x + y) ^ ((x ^ a) + (y ^ b))) & ((1 << n) - 1)) == g:
|
|
cnt += 1
|
|
return cnt / (1 << (2 * n))
|
|
|
|
|
|
def lm_prob(n, a, b, g):
|
|
if ((a ^ b ^ g) & 1):
|
|
return 0.0
|
|
w = 0
|
|
for i in range(n - 1):
|
|
ai, bi, gi = (a >> i) & 1, (b >> i) & 1, (g >> i) & 1
|
|
if ai == bi == gi:
|
|
if (((a >> (i + 1)) ^ (b >> (i + 1)) ^ (g >> (i + 1))) & 1) != ai:
|
|
return 0.0
|
|
else:
|
|
w += 1
|
|
return 2.0 ** -w
|
|
|
|
|
|
def cmd_selftest(args):
|
|
ok = True
|
|
# 1. evaluator against the Rust vectors
|
|
for name in sorted(os.listdir(args.vectors)):
|
|
if not name.endswith(".txt"):
|
|
continue
|
|
day, variant = name[:-4].split(".")
|
|
p = load_params(os.path.join(args.params_dir, f"{day}.{variant}.txt"))
|
|
n_ok = n_tot = 0
|
|
for line in open(os.path.join(args.vectors, name)):
|
|
cols = [[int(h, 16) for h in c.split()] for c in line.strip().split("|")]
|
|
s = cols[0]
|
|
for j in range(1, 5):
|
|
s = evaluate(s, ops_one_application(p, p["rk"][j - 1]))
|
|
n_tot += 1
|
|
n_ok += s == cols[j]
|
|
print(f"vectors {name}: {n_ok} of {n_tot} applications match")
|
|
ok &= n_ok == n_tot
|
|
# 2. the linear add automaton against brute force at n = 8 (random masks, and masks built to be valid)
|
|
rnd = random.Random(1)
|
|
n = 8
|
|
worst = 0.0
|
|
nz = 0
|
|
for _ in range(300):
|
|
u, v, w = rnd.randrange(256), rnd.randrange(256), rnd.randrange(256)
|
|
if rnd.random() < 0.5:
|
|
# a valid-by-construction mask: start from w and perturb a few bits
|
|
u = w ^ (1 << rnd.randrange(8)) if rnd.random() < 0.7 else w
|
|
v = w ^ (1 << rnd.randrange(8)) if rnd.random() < 0.7 else w
|
|
b = brute_lin_corr(n, u, v, w)
|
|
d = dp_lin_corr(n, u, v, w)
|
|
worst = max(worst, abs(b - d))
|
|
nz += b != 0
|
|
print(f"linear add automaton vs brute force n=8: 300 mask triples ({nz} nonzero), max |diff| {worst:.2e}")
|
|
ok &= worst < 1e-12
|
|
# 2b. the shifted-copy variant (low k0 bits of y zero): brute force with y's low bits forced to 0
|
|
worst = 0.0
|
|
for _ in range(200):
|
|
k0 = rnd.randrange(1, 4)
|
|
u, v, w = rnd.randrange(256), rnd.randrange(256), rnd.randrange(256)
|
|
if rnd.random() < 0.6:
|
|
u = w ^ (1 << rnd.randrange(8)) if rnd.random() < 0.7 else w
|
|
v = w ^ (1 << rnd.randrange(8)) if rnd.random() < 0.7 else w
|
|
tot = 0
|
|
for x in range(256):
|
|
for y in range(0, 256, 1 << k0):
|
|
z = (x + y) & 255
|
|
par = (bin(u & x).count("1") + bin(v & y).count("1") + bin(w & z).count("1")) & 1
|
|
tot += -1 if par else 1
|
|
b = tot / (256 * (256 >> k0))
|
|
d = dp_lin_corr(n, u, v, w, k0=k0)
|
|
worst = max(worst, abs(b - d))
|
|
print(f"linear add automaton with k0 low zero bits vs brute force n=8: 200 triples, max |diff| {worst:.2e}")
|
|
ok &= worst < 1e-12
|
|
# 3. Lipmaa-Moriai against brute force at n = 8
|
|
worst = 0.0
|
|
for _ in range(300):
|
|
a, b, g = rnd.randrange(256), rnd.randrange(256), rnd.randrange(256)
|
|
if rnd.random() < 0.6:
|
|
g = a ^ b ^ (rnd.randrange(256) & rnd.randrange(256) & 0xFE)
|
|
worst = max(worst, abs(brute_diff_prob(n, a, b, g) - lm_prob(n, a, b, g)))
|
|
print(f"Lipmaa-Moriai vs brute force n=8: 300 triples, max |diff| {worst:.2e}")
|
|
ok &= worst < 1e-12
|
|
# 3b. the SAT encodings against the rules they encode: add_lm against lm_prob, add_automaton against the
|
|
# automaton's best trail (max over sigma paths), on n = 32 with fixed random operands
|
|
from pysat.card import CardEnc as _CE
|
|
from pysat.solvers import Solver as _S
|
|
|
|
def sat_min_weight(cnf, weights, cap=40):
|
|
for bound in range(0, cap + 1):
|
|
card = _CE.atmost(lits=weights, bound=bound, top_id=cnf.nv)
|
|
with _S(name="cadical153", bootstrap_with=cnf.clauses + card.clauses) as s:
|
|
if s.solve():
|
|
return bound
|
|
return None
|
|
|
|
import math
|
|
bad = 0
|
|
for _ in range(40):
|
|
a, b = rnd.getrandbits(32), rnd.getrandbits(32)
|
|
g = a ^ b ^ (rnd.getrandbits(32) & rnd.getrandbits(32) & rnd.getrandbits(32) & 0x7FFFFFFE) if rnd.random() < 0.8 else rnd.getrandbits(32)
|
|
cnf = CNF()
|
|
m = DiffModel(cnf)
|
|
zs = m.add_lm(cnf.const_vec(a), cnf.const_vec(b))
|
|
for i in range(N):
|
|
cnf.add(zs[i] if (g >> i) & 1 else -zs[i])
|
|
got = sat_min_weight(cnf, m.weights)
|
|
pr = lm_prob(N, a, b, g)
|
|
want = None if pr == 0 else int(round(-math.log2(pr)))
|
|
bad += got != want
|
|
print(f"add_lm SAT encoding vs Lipmaa-Moriai rule n=32: 40 triples, mismatches {bad}")
|
|
ok &= bad == 0
|
|
|
|
def best_trail_weight(n, u, v, w, k0=0):
|
|
for i in range(k0):
|
|
if ((u >> i) & 1) != ((w >> i) & 1):
|
|
return None
|
|
states = {0: 0}
|
|
for i in range(k0, n):
|
|
a = ((u >> i) & 1) ^ ((w >> i) & 1)
|
|
b = ((v >> i) & 1) ^ ((w >> i) & 1)
|
|
wi = (w >> i) & 1
|
|
nxt = {}
|
|
for s_in, cost in states.items():
|
|
p = s_in ^ wi
|
|
for s_out in ([0] if i == n - 1 else [0, 1]):
|
|
if i == k0:
|
|
okk = (a == 0 and b == 0) if s_out == 0 else True
|
|
else:
|
|
wt = p + a + b
|
|
okk = (wt == 0) if s_out == 0 else (wt in (1, 3))
|
|
if okk:
|
|
c2 = cost + s_out
|
|
if nxt.get(s_out, 10**9) > c2:
|
|
nxt[s_out] = c2
|
|
states = nxt
|
|
return states.get(0)
|
|
|
|
bad = 0
|
|
for _ in range(40):
|
|
w = rnd.getrandbits(32)
|
|
u = w ^ (rnd.getrandbits(32) & rnd.getrandbits(32) & rnd.getrandbits(32))
|
|
v = w ^ (rnd.getrandbits(32) & rnd.getrandbits(32) & rnd.getrandbits(32))
|
|
k0 = rnd.choice([0, 0, 3, 7])
|
|
cnf = CNF()
|
|
m = LinModel(cnf)
|
|
uu, vv = m.add_automaton(cnf.const_vec(w), k0=k0)
|
|
for i in range(N):
|
|
cnf.add(uu[i] if (u >> i) & 1 else -uu[i])
|
|
if i >= k0:
|
|
cnf.add(vv[i] if (v >> i) & 1 else -vv[i])
|
|
got = sat_min_weight(cnf, m.weights)
|
|
want = best_trail_weight(N, u, v, w, k0=k0)
|
|
bad += got != want
|
|
print(f"add_automaton SAT encoding vs the automaton's best trail n=32: 40 triples, mismatches {bad}")
|
|
ok &= bad == 0
|
|
# 4. the multiply differential model: trail weight against sampled probability on the real word map
|
|
p = load_params(os.path.join(args.params_dir, "2026-10-03.real.txt"))
|
|
from pysat.solvers import Solver
|
|
for (din_desc, din) in [("MSB", 1 << 31), ("bit 0", 1), ("bit 30", 1 << 30), ("bits 31+5", (1 << 31) | 32)]:
|
|
i = 3
|
|
mulc, K = p["mul"][i], (p["rc"][i] + p["rk"][0]) & MASK32
|
|
cnf = CNF()
|
|
m = DiffModel(cnf)
|
|
ds = cnf.const_vec(din)
|
|
out = m.mul(ds, mulc)
|
|
# minimise weight: find the smallest bound with a model
|
|
from pysat.card import CardEnc
|
|
found = None
|
|
for bound in range(0, 40):
|
|
card = CardEnc.atmost(lits=m.weights, bound=bound, top_id=cnf.nv)
|
|
with Solver(name="cadical153", bootstrap_with=cnf.clauses + card.clauses) as s:
|
|
if s.solve():
|
|
ms = set(l for l in s.get_model() if l > 0)
|
|
found = (bound, lits_to_words(ms, [out])[0])
|
|
break
|
|
bound, dout = found
|
|
cnt = 0
|
|
S = 1 << 18
|
|
for _ in range(S):
|
|
x = rnd.getrandbits(32)
|
|
y0 = ((x ^ K) * mulc) & MASK32
|
|
y1 = (((x ^ din) ^ K) * mulc) & MASK32
|
|
cnt += (y0 ^ y1) == dout
|
|
import math
|
|
meas = -math.log2(cnt / S) if cnt else float("inf")
|
|
print(f"mult model word {i} din {din_desc}: best trail weight {bound} -> dout {dout:08x}; measured -log2 p = {meas:.2f} over 2^18 (trail <= differential expected)")
|
|
if din_desc == "MSB":
|
|
ok &= bound == 0 and cnt == S
|
|
# calibration: the model's cheapest trail from din (any output) against the sampled top transition
|
|
hist = {}
|
|
for _ in range(S):
|
|
x = rnd.getrandbits(32)
|
|
y0 = ((x ^ K) * mulc) & MASK32
|
|
y1 = (((x ^ din) ^ K) * mulc) & MASK32
|
|
d = y0 ^ y1
|
|
hist[d] = hist.get(d, 0) + 1
|
|
top = max(hist.values())
|
|
print(f" calibration din {din_desc}: model min weight (any output) {bound}; sampled top transition -log2 p = {-math.log2(top / S):.2f} ({top} of 2^18, {len(hist)} distinct outputs)")
|
|
# 5. the exact families: the msb family passes MSB-only words at cost 0 and refuses any other difference
|
|
cnf = CNF()
|
|
m = DiffModel(cnf, family="msb")
|
|
out = m.mul(cnf.const_vec(1 << 31), p["mul"][3])
|
|
with Solver(name="cadical153", bootstrap_with=cnf.clauses) as s:
|
|
r = s.solve()
|
|
ms = set(l for l in s.get_model() if l > 0) if r else set()
|
|
v = lits_to_words(ms, [out])[0] if r else None
|
|
print(f"msb family: MSB -> {v:08x} sat={r} weights={len(m.weights)}" if r else "msb family: MSB UNSAT")
|
|
ok &= r and v == (1 << 31) and len(m.weights) == 0
|
|
cnf = CNF()
|
|
m = DiffModel(cnf, family="msb")
|
|
m.mul(cnf.const_vec(1), p["mul"][3])
|
|
with Solver(name="cadical153", bootstrap_with=cnf.clauses) as s:
|
|
r = s.solve()
|
|
print(f"msb family: bit 0 sat={r} (must be False)")
|
|
ok &= not r
|
|
print("SELFTEST", "OK" if ok else "FAILED")
|
|
return 0 if ok else 1
|
|
|
|
|
|
def main():
|
|
ap = argparse.ArgumentParser()
|
|
sub = ap.add_subparsers(dest="cmd")
|
|
s = sub.add_parser("selftest")
|
|
s.add_argument("--vectors", required=True)
|
|
s.add_argument("--params-dir", required=True)
|
|
s = sub.add_parser("search")
|
|
s.add_argument("--kind", choices=["diff", "lin"], required=True)
|
|
s.add_argument("--family", default="general", help="diff: general|msb; lin: general|low2")
|
|
s.add_argument("--params", required=True)
|
|
s.add_argument("--apps", type=int, required=True)
|
|
s.add_argument("--state", required=True)
|
|
s.add_argument("--budget", type=float, default=1500)
|
|
s.add_argument("--cap", type=int, default=512)
|
|
s.add_argument("--solver", default="cadical195")
|
|
s.add_argument("--verifier", default=None, help="path of the attack-f2 binary: verify every found trail of weight <= --verify-max on the real code")
|
|
s.add_argument("--verify-max", type=int, default=12)
|
|
s.add_argument("--per-app-min", type=int, default=0, help="a proven lower bound on one application's trail weight (the k=1 result): every application of the chain is held to it (Matsui)")
|
|
s = sub.add_parser("show")
|
|
s.add_argument("--state", required=True)
|
|
args = ap.parse_args()
|
|
if args.cmd == "selftest":
|
|
sys.exit(cmd_selftest(args))
|
|
elif args.cmd == "search":
|
|
cmd_search(args)
|
|
elif args.cmd == "show":
|
|
print(show_line(json.load(open(args.state))))
|
|
else:
|
|
ap.print_help()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|