igneum/tools/attack/f2-mixer/model.py
2026-10-07 10:02:44 +00:00

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()