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