Bardia Taghavi · PhD Candidate, Florida Atlantic University (FAU) · PQSecure Technologies
This notebook is an interactive, number-driven companion to the workshop presentation. It executes the core mathematical and algorithmic mechanisms of ML-KEM (FIPS 203) and ML-DSA (FIPS 204), displaying the actual numbers, memory arrays, intermediate states, performance speedups, error boundaries, and rejection dynamics.
import sys, os, time, math, random, statistics, html, re
from hashlib import sha3_256, sha3_512, shake_128, shake_256
HAVE_DISPLAY = False
if "display" in globals() and "HTML" in globals():
HAVE_DISPLAY = True
else:
try:
import importlib
_disp = importlib.import_module("IPython.display")
HTML = getattr(_disp, "HTML")
display = getattr(_disp, "display")
HAVE_DISPLAY = True
except (ImportError, ModuleNotFoundError, AttributeError, Exception):
HAVE_DISPLAY = False
# ==============================================================================
# Catppuccin Editorial Palette & 4-Tier Styling Engine (from drawio standard)
# ==============================================================================
PALETTE = {
'secret': {'border': '#7287FD', 'bg': 'rgba(114, 135, 253, 0.12)', 'text': '#b4befe', 'name': 'Secret Key / Private Vector'},
'public': {'border': '#1E66F5', 'bg': 'rgba(30, 102, 245, 0.12)', 'text': '#89b4fa', 'name': 'Public Key / Verification Data'},
'input': {'border': '#DF8E1D', 'bg': 'rgba(223, 142, 29, 0.12)', 'text': '#f9e2af', 'name': 'Random Seed / Input Noise'},
'verified': {'border': '#40A02B', 'bg': 'rgba(64, 160, 43, 0.12)', 'text': '#a6e3a1', 'name': 'Verified / Valid Token'},
'token': {'border': '#04A5E5', 'bg': 'rgba(4, 165, 229, 0.12)', 'text': '#89dceb', 'name': 'Ciphertext / Intermediate Token'},
'reject': {'border': '#D20F39', 'bg': 'rgba(210, 15, 57, 0.12)', 'text': '#f38ba8', 'name': 'Rejection / Decryption Failure'},
'mauve': {'border': '#8839EF', 'bg': 'rgba(136, 57, 239, 0.12)', 'text': '#cba6f7', 'name': 'Decision / Performance Metric'},
'teal': {'border': '#179299', 'bg': 'rgba(23, 146, 153, 0.12)', 'text': '#94e2d5', 'name': 'Module / Core Algorithm'},
}
def badge(text, category='verified'):
theme = PALETTE.get(category, PALETTE['verified'])
return f'<span style="display:inline-block; background:{theme["bg"]}; color:{theme["border"]}; border:1px solid {theme["border"]}; border-radius:16px; padding:4px 14px; font-size:15px; font-weight:700; text-transform:uppercase; letter-spacing:0.7px;">{html.escape(str(text))}</span>'
def chip(val, highlight=None):
if highlight:
color, bg = highlight
elif isinstance(val, (int, float)):
if val < 0: color, bg = '#f38ba8', 'rgba(243,139,168,0.12)'
elif val == 0: color, bg = '#6c7086', 'rgba(108,112,134,0.12)'
else: color, bg = '#89b4fa', 'rgba(137,180,250,0.12)'
else:
color, bg = '#cdd6f4', 'rgba(255,255,255,0.06)'
return f'<span style="display:inline-block; background:{bg}; color:{color}; border:1px solid rgba(255,255,255,0.08); border-radius:6px; padding:3px 10px; margin:3px 3px; font-family:monospace; font-size:18px; font-weight:500;">{html.escape(str(val))}</span>'
def chips(values, max_items=12, highlight=None):
sub = list(values)[:max_items]
res = "".join(chip(v, highlight=highlight) for v in sub)
if len(values) > max_items:
res += f'<span style="color:#6c7086; font-size:15px; margin-left:8px;">... ({len(values)} total)</span>'
return res
def show_card(title, items, category='public', badge_text=None, footer=None):
if HAVE_DISPLAY:
theme = PALETTE.get(category, PALETTE['public'])
badge_html = f'<span style="background:{theme["bg"]}; color:{theme["border"]}; border:1px solid {theme["border"]}; border-radius:16px; padding:4px 14px; font-size:15px; font-weight:700; text-transform:uppercase; letter-spacing:0.7px; margin-left:14px;">{html.escape(str(badge_text))}</span>' if badge_text else ''
rows_html = ""
for k, v in items:
val_str = str(v) if ("<" in str(v) and ">" in str(v)) else f'<span style="font-family:monospace; color:#cdd6f4; font-size:18px;">{html.escape(str(v))}</span>'
rows_html += f'<div style="display:flex; justify-content:space-between; align-items:center; padding:10px 0; border-bottom:1px solid rgba(255,255,255,0.05); font-size:18px;"><span style="color:#a6adc8; font-weight:500; min-width:220px; margin-right:20px;">{html.escape(str(k))}</span><div style="text-align:right; flex-grow:1;">{val_str}</div></div>'
footer_html = f'<div style="margin-top:14px; padding-top:12px; border-top:1px solid rgba(255,255,255,0.08); color:#a6adc8; font-size:16px; font-style:italic;">{footer}</div>' if footer else ''
card_html = f'<div style="margin:18px 0; border-radius:12px; border:1px solid #313244; border-left:7px solid {theme["border"]}; background:#1e1e2e; padding:22px 28px; box-shadow:0 8px 20px rgba(0,0,0,0.3); font-family:sans-serif; line-height:1.5;"><div style="display:flex; align-items:center; margin-bottom:14px;"><span style="font-size:20px; font-weight:700; color:{theme["text"]};">{html.escape(str(title))}</span>{badge_html}</div><div>{rows_html}</div>{footer_html}</div>'
display(HTML(card_html))
else:
print(f"=== {title} ===")
for k, v in items:
clean_v = re.sub('<[^<]+?>', '', str(v))
print(f" {k}: {clean_v}")
if footer:
clean_f = re.sub('<[^<]+?>', '', str(footer))
print(f" {clean_f}")
def show_table(headers, rows, title=None, category='mauve', badge_text=None, highlight_col=None):
if HAVE_DISPLAY:
theme = PALETTE.get(category, PALETTE['mauve'])
badge_html = f'<span style="background:{theme["bg"]}; color:{theme["border"]}; border:1px solid {theme["border"]}; border-radius:16px; padding:4px 14px; font-size:15px; font-weight:700; text-transform:uppercase; letter-spacing:0.7px; margin-left:14px;">{html.escape(str(badge_text))}</span>' if badge_text else ''
header_title = f'<div style="display:flex; align-items:center; margin-bottom:14px;"><span style="font-size:20px; font-weight:700; color:{theme["text"]};">{html.escape(str(title))}</span>{badge_html}</div>' if title else ''
th_html = "".join(f'<th style="padding:14px 18px; font-size:16px; font-weight:600; text-transform:uppercase; letter-spacing:0.7px; color:#a6adc8;">{html.escape(str(h))}</th>' for h in headers)
tr_html = ""
for r_idx, row in enumerate(rows):
bg_row = 'background:rgba(255,255,255,0.025);' if r_idx % 2 == 1 else ''
td_html = ""
for c_idx, val in enumerate(row):
is_hl = (c_idx == highlight_col)
val_str = str(val)
cell_content = val_str if ("<" in val_str and ">" in val_str) else f'<span style="font-family:monospace; font-size:18px; color:{theme["text"] if is_hl else "#cdd6f4"}; font-weight:{600 if is_hl else 400};">{html.escape(val_str)}</span>'
td_html += f'<td style="padding:12px 18px; border-bottom:1px solid rgba(255,255,255,0.04); font-size:18px;">{cell_content}</td>'
tr_html += f'<tr style="{bg_row}">{td_html}</tr>'
table_html = f'<div style="margin:18px 0; border-radius:12px; border:1px solid #313244; border-left:7px solid {theme["border"]}; background:#1e1e2e; padding:22px 28px; box-shadow:0 8px 20px rgba(0,0,0,0.3); font-family:sans-serif; line-height:1.5;">{header_title}<div style="overflow-x:auto;"><table style="width:100%; border-collapse:collapse; margin:4px 0; text-align:left; font-size:18px;"><thead><tr style="border-bottom:2px solid #313244; background:#181825;">{th_html}</tr></thead><tbody>{tr_html}</tbody></table></div></div>'
display(HTML(table_html))
else:
if title: print(f"\n{title}:")
table(headers, rows)
def table(headers, rows):
widths = [len(str(h)) for h in headers]
for row in rows:
for i, val in enumerate(row):
val_clean = re.sub('<[^<]+?>', '', str(val))
widths[i] = max(widths[i], len(val_clean))
fmt = " ".join("%-" + str(w) + "s" for w in widths)
print(fmt % tuple(headers))
print(" ".join("-" * w for w in widths))
for row in rows:
clean_row = [re.sub('<[^<]+?>', '', str(v)) for v in row]
print(fmt % tuple(clean_row))
show_card("PQC Embedded Demonstration Environment",
[("Runtime Engine", f"Python {sys.version.split()[0]}"),
("Cryptographic Hash Primitives", "FIPS 202 Keccak (SHA3-256, SHA3-512, SHAKE128, SHAKE256)"),
("Display Theme", "Catppuccin Misto / Latte 4-Tier Editorial Cards"),
("Execution Mode", "100% Pure Python + Zero Compiled Dependencies")],
category='teal', badge_text="Initialized")
PASSED (32.4 ms)
---
The foundation of modern lattice cryptography is noisy linear algebra over polynomial rings.
Instead of abstract geometry, this section demonstrates how coordinates, rings, and error boundaries work with concrete numbers.
# 1.1 The Closest Vector Problem (CVP) with Concrete Numbers
def round_to_lattice(target, b1, b2):
det = b1[0] * b2[1] - b1[1] * b2[0]
c1 = (target[0] * b2[1] - target[1] * b2[0]) / det
c2 = (b1[0] * target[1] - b1[1] * target[0]) / det
k1, k2 = round(c1), round(c2)
return (k1 * b1[0] + k2 * b2[0], k1 * b1[1] + k2 * b2[1]), (k1, k2)
target = (2.7, 3.4)
good_b1, good_b2 = (1, 0), (0, 1)
bad_b1, bad_b2 = (10, 11), (11, 12)
landed_good, coords_good = round_to_lattice(target, good_b1, good_b2)
landed_bad, coords_bad = round_to_lattice(target, bad_b1, bad_b2)
dist_good = math.hypot(target[0] - landed_good[0], target[1] - landed_good[1])
dist_bad = math.hypot(target[0] - landed_bad[0], target[1] - landed_bad[1])
show_card("Closest Vector Problem: Babai Rounding Comparison",
[("Continuous Target Point", chip(target)),
("Good Basis (Orthogonal)", f"b1={good_b1}, b2={good_b2}"),
(" -> Decoded Integer Point", f"{landed_good} (combo: {coords_good})"),
(" -> Distance to Target", f"{dist_good:.3f} " + badge("Exact Closest Point", "verified")),
("Bad Basis (Skewed)", f"b1={bad_b1}, b2={bad_b2}"),
(" -> Decoded Integer Point", f"{landed_bad} (combo: {coords_bad})"),
(" -> Distance to Target", f"{dist_bad:.3f} " + badge(f"Error {dist_bad/dist_good:.1f}x Larger", "reject"))],
category='teal', badge_text="CVP Decoding",
footer="Takeaway: The public key is a 'bad basis' (hard CVP). The private key is a 'good basis' (easy noise removal).")
PASSED (0.2 ms)
Everything in ML-KEM and ML-DSA happens modulo a prime q.
Both standards mean the centred representative in (-q/2, q/2] whenever they state that a value is "small" (noise, secrets, errors).
def modpm(r, a):
r %= a
return r - a if r > a // 2 else r
Q_KEM = 3329 # ML-KEM (12-bit prime)
Q_DSA = 8380417 # ML-DSA = 2**23 - 2**13 + 1 (23-bit prime)
mod_rows = []
for x in (10, 1664, 1665, 3328, 0, 7):
std = x % 13
cnt = modpm(x, 13)
cnt_styled = f'<span style="color:{"#f38ba8" if cnt < 0 else ("#6c7086" if cnt == 0 else "#89b4fa")}; font-weight:600;">{cnt:+d}</span>'
mod_rows.append([str(x), str(std), cnt_styled])
show_table(["Value x", "Standard mod 13", "Centred mod 13 in (-6, +6]"],
mod_rows, title="Standard vs Centred Modulo Reduction (q = 13)", category='public')
s = 12345
show_card("Hardware Modular Arithmetic Properties",
[("ML-KEM Modulus", f"q = {Q_KEM} = 2^8 * 13 + 1 ({Q_KEM.bit_length()} bits)"),
("ML-DSA Modulus", f"q = {Q_DSA} = 2^23 - 2^13 + 1 ({Q_DSA.bit_length()} bits)"),
("ML-DSA Zero-DSP Trick", f"s * q = (s << 23) - (s << 13) + s"),
("Verification for s=12345", f"s*q = {s*Q_DSA} == {(s<<23) - (s<<13) + s} " + badge("Verified: 2 shifts, 1 sub, 1 add", "verified"))],
category='teal', badge_text="Zero DSP Multipliers")
PASSED (0.2 ms)
| Value x | Standard mod 13 | Centred mod 13 in (-6, +6] |
|---|---|---|
| 10 | 10 | -3 |
| 1664 | 0 | +0 |
| 1665 | 1 | +1 |
| 3328 | 0 | +0 |
| 0 | 0 | +0 |
| 7 | 7 | -6 |
A polynomial is an array of n coefficients modulo q.
Multiplication in R_q = Z_q[X]/(X^n + 1) wraps around at degree n with a sign flip (X^n = -1).
def poly_mul_schoolbook(a, b, q, n=None):
n = n or len(a)
out = [0] * n
for i, ai in enumerate(a):
if not ai: continue
for j, bj in enumerate(b):
k = i + j
if k < n:
out[k] = (out[k] + ai * bj) % q
else:
out[k - n] = (out[k - n] - ai * bj) % q
return out
a = [1, 2, 3, 4]
b = [5, 6, 7, 8]
c = poly_mul_schoolbook(a, b, 17)
cyc = [0] * 4
for i in range(4):
for j in range(4):
cyc[(i + j) % 4] = (cyc[(i + j) % 4] + a[i] * b[j]) % 17
show_card("Negacyclic Polynomial Multiplication (n = 4, q = 17)",
[("Input Polynomial a", chips(a)),
("Input Polynomial b", chips(b)),
("Negacyclic Product a * b mod (X^4 + 1)", chips(c) + " " + badge("Correct Ring Product", "verified")),
("Plain Cyclic Wrap (without sign flip)", chips(cyc) + " " + badge("Wrong (Missing X^n = -1)", "reject")),
("Constant Term c[0] Derivation", "(1*5 - 2*8 - 3*7 - 4*6) mod 17 = -56 mod 17 = 12")],
category='mauve', badge_text="Negacyclic Wrap")
PASSED (0.2 ms)
Now the actual mechanism behind ML-KEM, at a size you can print: n = 8 coefficients, modulus q = 97, noise bound eta = 1.
One message bit is stretched to q // 2 = 48, buried under secret-dependent noise, transmitted as ciphertext (u, v), and recovered by rounding.
def toy_keygen(q, n, eta, rng):
a = [rng.randrange(q) for _ in range(n)]
s = [rng.randrange(-eta, eta + 1) % q for _ in range(n)]
e = [rng.randrange(-eta, eta + 1) % q for _ in range(n)]
t = [(x + y) % q for x, y in zip(poly_mul_schoolbook(a, s, q), e)]
return (a, t), s
def toy_encrypt(pk, bit, q, n, eta, rng):
a, t = pk
y = [rng.randrange(-eta, eta + 1) % q for _ in range(n)]
e1 = [rng.randrange(-eta, eta + 1) % q for _ in range(n)]
e2 = rng.randrange(-eta, eta + 1) % q
u = [(x + y_) % q for x, y_ in zip(poly_mul_schoolbook(a, y, q), e1)]
v = (poly_mul_schoolbook(t, y, q)[0] + e2 + bit * (q // 2)) % q
return u, v, y, e1, e2
def toy_decrypt(sk, ct, q):
u, v = ct
su = poly_mul_schoolbook(sk, u, q)[0]
w = (v - su) % q
bit_rec = 0 if abs(modpm(w, q)) < q // 4 else 1
return bit_rec, w, su
Q, Nn, ETA = 97, 8, 1
rng = random.Random(7)
pk, sk = toy_keygen(Q, Nn, ETA, rng)
a, t = pk
show_card(f"Toy KEM Keys & Parameters (q = {Q}, n = {Nn}, eta = {ETA})",
[("Public Matrix a", chips(a)),
("Public Vector t = a*s + e", chips(t)),
("Secret Key s (centred in [-1, +1])", chips([modpm(x, Q) for x in sk]))],
category='secret', badge_text="Key Generation")
for bit in (0, 1):
u, v, y, e1, e2 = toy_encrypt(pk, bit, Q, Nn, ETA, rng)
got, w, su = toy_decrypt(sk, (u, v), Q)
target_offset = 0 if bit == 0 else Q // 2
cnt_w = modpm(w, Q)
w_cond = abs(cnt_w) < Q // 4
status_badge = badge("Success: Bit Recovered", "verified") if got == bit else badge("Failed", "reject")
show_card(f"Toy KEM Transmission [Bit = {bit}]",
[("Ephemeral Secret y", chips([modpm(x, Q) for x in y])),
("Injected Noise (e1, e2)", f"e1={chips([modpm(x, Q) for x in e1])}, e2={chip(modpm(e2, Q))}"),
("Ciphertext u (8 ring coefficients)", chips(u)),
("Ciphertext v (scalar mod 97)", chip(v)),
("Decryption Product s*u[0] mod 97", chip(su)),
("Noisy Message w = v - s*u", f"{w} mod 97 (centred: {chip(cnt_w)})"),
("Target Center (bit * 48)", f"{target_offset}"),
("Decision Boundary (|w| < 24)", f"|{cnt_w:+d}| < 24 is {w_cond}"),
("Recovered Bit", f"Bit {got} " + status_badge)],
category='token', badge_text=f"Bit {bit} Transmission")
PASSED (0.5 ms)
Security requires noise (eta) to hide the secret from lattice reduction attacks.
Correctness requires noise to remain within the decoding boundary (< q/4).
Watch what happens to the actual numbers when noise eta exceeds the safe margin.
# 1. Concrete Decryption Failure Trace (Watching an actual bit flip caused by excessive noise)
rng_err = random.Random(42)
for trial in range(1, 100):
pk_t, sk_t = toy_keygen(Q, Nn, 4, rng_err)
bit_sent = 0
u_err, v_err, _, _, _ = toy_encrypt(pk_t, bit_sent, Q, Nn, 4, rng_err)
bit_rec, w_err, su_err = toy_decrypt(sk_t, (u_err, v_err), Q)
if bit_rec != bit_sent:
cnt_w_err = modpm(w_err, Q)
show_card(f"Concrete Decryption Failure (Trial #{trial}, eta = 4)",
[("Sent Message Bit", chip(bit_sent)),
("Ciphertext u", chips(u_err)),
("Ciphertext v", chip(v_err)),
("Reconstructed Noisy w", f"{w_err} mod 97 (centred: {chip(cnt_w_err)})"),
("Decision Boundary (+/- q/4)", "+/- 24"),
("Overflow Check", f"|{cnt_w_err:+d}| >= 24 -> Bit flips from {bit_sent} to {bit_rec}!"),
("Decryption Verdict", badge("Bit Flip Error (Decoding Failure)", "reject"))],
category='reject', badge_text="Noise Overflow")
break
# 2. Failure rate measurement across noise bounds
def failure_rate(q, n, eta, trials=1000, seed=0):
rng_f = random.Random(seed)
bad = 0
for _ in range(trials):
pk_f, sk_f = toy_keygen(q, n, eta, rng_f)
bit = rng_f.randrange(2)
u, v, _, _, _ = toy_encrypt(pk_f, bit, q, n, eta, rng_f)
got, _, _ = toy_decrypt(sk_f, (u, v), q)
if got != bit:
bad += 1
return bad / trials
rows_fail = []
for e in (1, 2, 3, 4, 6, 8, 10):
rate = failure_rate(97, 8, e)
if rate == 0:
st = badge("Zero Failures (Safe)", "verified")
elif rate < 0.15:
st = badge(f"Minor Errors ({100*rate:.1f}%)", "input")
else:
st = badge(f"Critical Failure ({100*rate:.1f}%)", "reject")
rows_fail.append([f"eta = {e}", f"{100 * rate:5.1f} %", st])
show_table(["Noise Bound (eta)", "Empirical Failure Rate", "Status"],
rows_fail, title="Failure Rate vs Noise Bound (1000 trials each, q = 97, n = 8)", category='reject')
PASSED (141.7 ms)
| Noise Bound (eta) | Empirical Failure Rate | Status |
|---|---|---|
| eta = 1 | 0.0 % | Zero Failures (Safe) |
| eta = 2 | 0.1 % | Minor Errors (0.1%) |
| eta = 3 | 11.5 % | Minor Errors (11.5%) |
| eta = 4 | 34.0 % | Critical Failure (34.0%) |
| eta = 6 | 49.8 % | Critical Failure (49.8%) |
| eta = 8 | 50.8 % | Critical Failure (50.8%) |
| eta = 10 | 49.8 % | Critical Failure (49.8%) |
def brute_force_secret(pk, q, n, eta):
a, t = pk
tried = 0
for combo in range(0, (2 * eta + 1) ** n):
s, rest = [], combo
for _ in range(n):
s.append((rest % (2 * eta + 1)) - eta)
rest //= (2 * eta + 1)
tried += 1
guess = poly_mul_schoolbook(a, [x % q for x in s], q)
if all(abs(modpm(t[i] - guess[i], q)) <= eta for i in range(n)):
return [x % q for x in s], tried
return None, tried
rng_bf = random.Random(11)
pk_bf, sk_bf = toy_keygen(97, 4, 1, rng_bf)
found, tried = brute_force_secret(pk_bf, 97, 4, 1)
show_card("Lattice Security & Brute-Force Feasibility",
[("Toy Parameters (n = 4, eta = 1)", f"Search space: 3^4 = 81 candidates"),
("Brute-Force Attack Result", f"Found secret in {tried} attempts " + badge("Trivially Broken", "reject")),
("ML-KEM-768 Secret Dimension", "k * n = 3 * 256 = 768 coefficients in [-2, +2]"),
("ML-KEM-768 Search Space", "5^768 ≈ 10^536 candidates"),
("Physical Comparison", "Total atoms in observable universe ≈ 10^80 " + badge("Information-Theoretically Impenetrable", "verified"))],
category='teal', badge_text="Search Space Scaling")
PASSED (0.3 ms)
---
Schoolbook polynomial multiplication costs O(n^2) = 65,536 modular multiplies for n = 256.
The NTT computes it in O(n log n) ≈ 3,300 modular multiplies by moving to the frequency domain.
This section shows:
1. INTT(NTT(a)) == a verified coefficient-by-coefficient with actual numbers.
2. The actual numbers in memory for ML-KEM (128 pairs of values) and ML-DSA (256 scalar evaluations).
3. INTT(pointwise(NTT(a), NTT(b))) == poly_mul_schoolbook(a, b) with measured Python speedups and hardware arithmetic savings.
def find_root_of_unity(q, order):
if (q - 1) % order: return None
for g in range(2, q):
if pow(g, order, q) == 1:
if all(pow(g, order // p, q) != 1 for p in (2, 3, 5, 7, 11, 13) if not order % p):
return g
return None
def factorise(n):
factors, d = {}, 2
while d * d <= n:
while not n % d:
factors[d] = factors.get(d, 0) + 1
n //= d
d += 1
if n > 1:
factors[n] = factors.get(n, 0) + 1
return factors
root_rows = []
for name, q in (("ML-KEM", Q_KEM), ("ML-DSA", Q_DSA)):
f_str = " * ".join(f"{p}^{e}" for p, e in factorise(q - 1).items())
has_256 = find_root_of_unity(q, 256) is not None
has_512 = find_root_of_unity(q, 512) is not None
root_rows.append([name, str(q), f_str,
badge("Exists", "verified") if has_256 else badge("None", "reject"),
badge("Exists", "verified") if has_512 else badge("None", "reject")])
show_table(["Standard", "Modulus q", "Prime Factorisation of (q - 1)", "256th Root", "512th Root"],
root_rows, title="Roots of Unity and Modulus Factorisation", category='mauve')
PASSED (26.6 ms)
| Standard | Modulus q | Prime Factorisation of (q - 1) | 256th Root | 512th Root |
|---|---|---|---|---|
| ML-KEM | 3329 | 2^8 * 13^1 | Exists | None |
| ML-DSA | 8380417 | 2^13 * 3^1 * 11^1 * 31^1 | Exists | Exists |
def bitrev(i, bits):
return int(format(i, "0%db" % bits)[::-1], 2)
def make_zetas(q, zeta, count, bits):
return [pow(zeta, bitrev(i, bits), q) for i in range(count)]
ZETA_KEM = 17 # 256th root of unity mod 3329
ZETA_DSA = 1753 # 512th root of unity mod 8380417
ZETAS_KEM = make_zetas(Q_KEM, ZETA_KEM, 128, 7)
ZETAS_DSA = make_zetas(Q_DSA, ZETA_DSA, 256, 8)
GAMMAS_KEM = [pow(ZETA_KEM, 2 * bitrev(i, 7) + 1, Q_KEM) for i in range(128)]
NINV_KEM = pow(128, Q_KEM - 2, Q_KEM)
NINV_DSA = pow(256, Q_DSA - 2, Q_DSA)
show_card("NTT Twiddle Factors & Architectures",
[("ML-KEM Generator (order 256)", f"zeta = {ZETA_KEM}, n^-1 = {NINV_KEM}"),
(" First 8 Bit-Reversed Twiddles", chips(ZETAS_KEM[:8])),
("ML-DSA Generator (order 512)", f"zeta = {ZETA_DSA}, n^-1 = {NINV_DSA}"),
(" First 8 Bit-Reversed Twiddles", chips(ZETAS_DSA[:8]))],
category='teal', badge_text="Twiddle Tables",
footer="ML-KEM stops at 7 stages (128 degree-1 pairs); ML-DSA runs 8 full stages (256 scalar values).")
PASSED (0.4 ms)
def ntt_generic(f, q, zetas, stages):
a = list(f)
k = 1
len_ = 128
for stage in range(stages):
start = 0
while start < 256:
zeta = zetas[k]; k += 1
for j in range(start, start + len_):
t = (zeta * a[j + len_]) % q
a[j + len_] = (a[j] - t) % q
a[j] = (a[j] + t) % q
start = start + 2 * len_
len_ //= 2
return a
def INTT_generic(f, q, zetas, stages, ninv):
a = list(f)
k = len(zetas) - 1
len_ = 256 // (2 ** stages)
for stage in range(stages):
start = 0
while start < 256:
zeta = zetas[k]; k -= 1
for j in range(start, start + len_):
t = a[j]
a[j] = (t + a[j + len_]) % q
a[j + len_] = (zeta * (a[j + len_] - t)) % q
start = start + 2 * len_
len_ *= 2
return [(x * ninv) % q for x in a]
def ntt_kem(f): return ntt_generic(f, Q_KEM, ZETAS_KEM, 7)
def INTT_kem(f): return INTT_generic(f, Q_KEM, ZETAS_KEM, 7, NINV_KEM)
def ntt_dsa(f): return ntt_generic(f, Q_DSA, ZETAS_DSA, 8)
def INTT_dsa(f): return INTT_generic(f, Q_DSA, ZETAS_DSA, 8, NINV_DSA)
def base_case_multiply(a0, a1, b0, b1, gamma, q=Q_KEM):
return ((a0 * b0 + a1 * b1 % q * gamma) % q, (a0 * b1 + a1 * b0) % q)
def pointwise_kem(f, g):
h = [0] * 256
for i in range(128):
h[2*i], h[2*i+1] = base_case_multiply(f[2*i], f[2*i+1],
g[2*i], g[2*i+1], GAMMAS_KEM[i])
return h
def pointwise_dsa(f, g):
return [x * y % Q_DSA for x, y in zip(f, g)]
# Demonstration 1: ML-KEM NTT Verification
a_kem = [(i * 13 + 7) % Q_KEM for i in range(256)]
a_hat_kem = ntt_kem(a_kem)
a_rec_kem = INTT_kem(a_hat_kem)
diff_kem = max(abs(x - y) for x, y in zip(a_kem, a_rec_kem))
show_card("ML-KEM NTT Verification: INTT(NTT(a)) == a",
[("Input Polynomial a[:12]", chips(a_kem[:12])),
("NTT Transformed a_hat[:12]", chips(a_hat_kem[:12])),
("Reconstructed INTT(NTT(a))[:12]", chips(a_rec_kem[:12])),
("Exact Inversion Check", f"All 256 coefficients match (Max diff: {diff_kem}) " + badge("Verified: Exact Inversion", "verified"))],
category='teal', badge_text="ML-KEM Inversion")
mem_rows = []
for i in (0, 1, 2, 3, 127):
mem_rows.append([f"Block {i:3d}", f"indices ({2*i:3d}, {2*i+1:3d})",
f"({a_hat_kem[2*i]:4d}, {a_hat_kem[2*i+1]:4d})", f"mod (X^2 - {GAMMAS_KEM[i]:4d})"])
show_table(["Block Number", "Memory Indices", "Stored Value Pair [a0, a1]", "Modulus Ideal"],
mem_rows, title="ML-KEM In-Memory Layout: 128 Pairs of Degree-1 Polynomials", category='public')
# Demonstration 2: ML-DSA NTT Verification
a_dsa = [(i * 131 + 17) % Q_DSA for i in range(256)]
a_hat_dsa = ntt_dsa(a_dsa)
a_rec_dsa = INTT_dsa(a_hat_dsa)
diff_dsa = max(abs(x - y) for x, y in zip(a_dsa, a_rec_dsa))
show_card("ML-DSA NTT Verification: Full 8-Stage 256-Point Inversion",
[("Input Polynomial a[:8]", chips(a_dsa[:8])),
("NTT Transformed a_hat[:8]", chips(a_hat_dsa[:8])),
("Reconstructed INTT(NTT(a))[:8]", chips(a_rec_dsa[:8])),
("Exact Inversion Check", f"All 256 coefficients match (Max diff: {diff_dsa}) " + badge("Verified: Exact Inversion", "verified"))],
category='teal', badge_text="ML-DSA Inversion")
PASSED (1.0 ms)
| Block Number | Memory Indices | Stored Value Pair [a0, a1] | Modulus Ideal |
|---|---|---|---|
| Block 0 | indices ( 0, 1) | (1199, 3278) | mod (X^2 - 17) |
| Block 1 | indices ( 2, 3) | (1457, 2938) | mod (X^2 - 3312) |
| Block 2 | indices ( 4, 5) | ( 708, 749) | mod (X^2 - 2761) |
| Block 3 | indices ( 6, 7) | ( 632, 2910) | mod (X^2 - 568) |
| Block 127 | indices (254, 255) | (2462, 409) | mod (X^2 - 1175) |
rng_mul = random.Random(42)
a_poly = [rng_mul.randrange(Q_KEM) for _ in range(256)]
b_poly = [rng_mul.randrange(Q_KEM) for _ in range(256)]
t0 = time.perf_counter()
for _ in range(15):
c_school = poly_mul_schoolbook(a_poly, b_poly, Q_KEM)
t_school = (time.perf_counter() - t0) / 15
t0 = time.perf_counter()
for _ in range(15):
c_ntt = INTT_kem(pointwise_kem(ntt_kem(a_poly), ntt_kem(b_poly)))
t_ntt = (time.perf_counter() - t0) / 15
diff_mul = max(abs(x - y) for x, y in zip(c_school, c_ntt))
show_card("Polynomial Multiplication: Schoolbook vs. Fast NTT",
[("Input a[:8]", chips(a_poly[:8])),
("Input b[:8]", chips(b_poly[:8])),
("Schoolbook Product c[:8]", chips(c_school[:8])),
("Fast NTT Product c[:8]", chips(c_ntt[:8])),
("Equivalence Check", f"Identical for all 256 coefficients (Max diff: {diff_mul}) " + badge("Exact Match", "verified")),
("Execution Speedup (Python)", f"Schoolbook: {t_school*1000:.2f} ms | NTT: {t_ntt*1000:.2f} ms -> " + badge(f"{t_school/t_ntt:.1f}x Speedup", "mauve"))],
category='mauve', badge_text="Arithmetic Equivalence")
show_table(["Multiplication Algorithm", "Modular Multiplies", "Time Complexity", "Hardware Workload"],
[["Schoolbook Convolution", "256 x 256 = 65,536", "O(n^2)", "Baseline (100%)"],
["Fast NTT (ML-KEM)", "2*896 (fwd) + 640 (ptwise) + 896 (inv) = 3,328", "O(n log n)", badge("95% Workload Reduction", "verified")],
["Fast NTT (ML-DSA)", "2*1024 (fwd) + 256 (ptwise) + 1024 (inv) = 3,328", "O(n log n)", badge("95% Workload Reduction", "verified")]],
title="Hardware Multiplier Count: Schoolbook vs Fast NTT", category='mauve')
PASSED (53.5 ms)
| Multiplication Algorithm | Modular Multiplies | Time Complexity | Hardware Workload |
|---|---|---|---|
| Schoolbook Convolution | 256 x 256 = 65,536 | O(n^2) | Baseline (100%) |
| Fast NTT (ML-KEM) | 2*896 (fwd) + 640 (ptwise) + 896 (inv) = 3,328 | O(n log n) | 95% Workload Reduction |
| Fast NTT (ML-DSA) | 2*1024 (fwd) + 256 (ptwise) + 1024 (inv) = 3,328 | O(n log n) | 95% Workload Reduction |
q_toy = 97
omega_toy = 36
twiddles = [pow(omega_toy, bitrev(k, 3), q_toy) for k in range(8)]
reg = [1, 2, 3, 4, 5, 6, 7, 8]
trace_rows = [["Stage 0 (Input)", chips(reg)]]
for i in range(4):
u = reg[i]
v = (reg[i+4] * twiddles[1]) % q_toy
reg[i] = (u + v) % q_toy
reg[i+4] = (u - v) % q_toy
trace_rows.append(["Stage 1 (len=4)", chips(reg)])
for block in (0, 4):
for i in range(2):
idx = block + i
u = reg[idx]
v = (reg[idx+2] * twiddles[2 + (block // 4)]) % q_toy
reg[idx] = (u + v) % q_toy
reg[idx+2] = (u - v) % q_toy
trace_rows.append(["Stage 2 (len=2)", chips(reg)])
for block in (0, 2, 4, 6):
u = reg[block]
v = (reg[block+1] * twiddles[4 + (block // 2)]) % q_toy
reg[block] = (u + v) % q_toy
reg[block+1] = (u - v) % q_toy
trace_rows.append(["Stage 3 (Final Output)", chips(reg)])
show_table(["Cooley-Tukey Pipeline Stage", "Hardware Register Bank Contents"],
trace_rows, title="8-Point Cooley-Tukey Butterfly Hardware Register Trace", category='teal')
PASSED (0.2 ms)
| Cooley-Tukey Pipeline Stage | Hardware Register Bank Contents |
|---|---|
| Stage 0 (Input) | 12345678 |
| Stage 1 (len=4) | 15774278424611 |
| Stage 2 (len=2) | 303102648252323 |
| Stage 3 (Final Output) | 7978603723737568 |
A butterfly is one modular multiply plus a modular add and subtract. The multiply is easy; the *reduction* is where the gates go. Barrett reduction replaces the division by a multiply-by-a-constant, and for ML-DSA's q even that constant multiply collapses into shifts.
def barrett_setup(q):
t = q.bit_length()
return t, (1 << (2 * t)) // q
def barrett_reduce(x, q, t, mu):
s = (x * mu) >> (2 * t)
r = x - s * q
while r >= q:
r -= q
return r
barr_rows = []
for q, name in ((Q_KEM, "ML-KEM"), (Q_DSA, "ML-DSA")):
t, mu = barrett_setup(q)
rng2 = random.Random(q)
worst = 0
for _ in range(10000):
x = rng2.randrange(q * q)
r = barrett_reduce(x, q, t, mu)
s = (x * mu) >> (2 * t)
worst = max(worst, (x - s * q) // q)
barr_rows.append([name, str(q), f"{t} bits", f"floor(2^{2*t} / q) = {mu}",
f"Worst corrections: {worst} " + badge("10k Passed", "verified")])
show_table(["Standard", "Modulus q", "Bit Length t", "Precomputed Multiplier mu", "Barrett Validation"],
barr_rows, title="Barrett Constant-Multiplier Reduction Validation", category='mauve')
PASSED (8.3 ms)
| Standard | Modulus q | Bit Length t | Precomputed Multiplier mu | Barrett Validation |
|---|---|---|---|---|
| ML-KEM | 3329 | 12 bits | floor(2^24 / q) = 5039 | Worst corrections: 1 10k Passed |
| ML-DSA | 8380417 | 23 bits | floor(2^46 / q) = 8396807 | Worst corrections: 1 10k Passed |
---
Complete, standard-compliant implementation of ML-KEM (Module-LWE Key Encapsulation Mechanism).
All internal objects (s, e, t, y, e1, e2, u, v, w) are inspected as concrete arrays of numbers.
# ---- Keccak, with a counter so section 5 can see where the time goes -------
KECCAK = {"calls": 0, "perms": 0, "detail": {}}
RATE = {"SHAKE128": 168, "SHAKE256": 136, "SHA3-256": 136, "SHA3-512": 72}
def _count(kind, in_len, out_len):
KECCAK["calls"] += 1
rate = RATE[kind]
perms = -(-(in_len + 1) // rate) + max(0, -(-out_len // rate) - 1)
KECCAK["perms"] += perms
d = KECCAK["detail"].setdefault(kind, [0, 0])
d[0] += 1
d[1] += perms
return perms
def keccak_reset():
KECCAK["calls"] = 0
KECCAK["perms"] = 0
KECCAK["detail"] = {}
def _shake128(data, out_len):
_count("SHAKE128", len(data), out_len)
return shake_128(data).digest(out_len)
def _shake256(data, out_len):
_count("SHAKE256", len(data), out_len)
return shake_256(data).digest(out_len)
def _sha3_256(data):
_count("SHA3-256", len(data), 32)
return sha3_256(data).digest()
def _sha3_512(data):
_count("SHA3-512", len(data), 64)
return sha3_512(data).digest()
class Squeeze:
"""A SHAKE stream we can pull bytes from a block at a time."""
def __init__(self, kind, seed, block=504):
self.kind, self.seed, self.block = kind, seed, block
self.buf, self.pos = b"", 0
def take(self, n):
while self.pos + n > len(self.buf):
want = len(self.buf) + self.block
fn = _shake128 if self.kind == "SHAKE128" else _shake256
# Count only the extra squeeze blocks, not a whole re-absorption.
if self.buf:
_count(self.kind, 0, self.block)
self.buf = (shake_128 if self.kind == "SHAKE128"
else shake_256)(self.seed).digest(want)
else:
self.buf = fn(self.seed, want)
out = self.buf[self.pos:self.pos + n]
self.pos += n
return out
# ---- ML-KEM parameters ----------------------------------------------------
KEM_PARAMS = {
512: dict(k=2, eta1=3, eta2=2, du=10, dv=4),
768: dict(k=3, eta1=2, eta2=2, du=10, dv=4),
1024: dict(k=4, eta1=2, eta2=2, du=11, dv=5),
}
def kem_H(d): return _sha3_256(d)
def kem_G(d): out = _sha3_512(d); return out[:32], out[32:]
def kem_J(d): return _shake256(d, 32)
def kem_prf(eta, s, b): return _shake256(s + bytes([b]), 64 * eta)
# ---- sampling -------------------------------------------------------------
def sample_ntt(seed):
"""Algorithm 7: uniform mod q, by rejecting 12-bit values >= q."""
st = Squeeze("SHAKE128", seed)
out = []
while len(out) < 256:
b = st.take(3)
d1 = b[0] + 256 * (b[1] % 16)
d2 = (b[1] // 16) + 16 * b[2]
if d1 < Q_KEM:
out.append(d1)
if d2 < Q_KEM and len(out) < 256:
out.append(d2)
return out
def sample_poly_cbd(eta, data):
"""Algorithm 8: centred binomial. Count bits, subtract. No rejection."""
bits = [(data[i // 8] >> (i % 8)) & 1 for i in range(8 * len(data))]
out = []
for i in range(256):
x = sum(bits[2 * i * eta + j] for j in range(eta))
y = sum(bits[2 * i * eta + eta + j] for j in range(eta))
out.append((x - y) % Q_KEM)
return out
print("Keccak helpers and samplers ready.")
Keccak helpers and samplers ready. PASSED (0.5 ms)
# ---- packing and compression ---------------------------------------------
def byte_encode(d, f):
"""Algorithm 5: 256 values of d bits -> 32d bytes."""
acc = 0
for i, v in enumerate(f):
acc |= (v & ((1 << d) - 1)) << (d * i)
return acc.to_bytes(32 * d, "little")
def byte_decode(d, b):
"""Algorithm 6."""
acc = int.from_bytes(b, "little")
mask = (1 << d) - 1
out = [(acc >> (d * i)) & mask for i in range(256)]
return [v % Q_KEM for v in out] if d == 12 else out
def compress(d, x):
return [(((v << d) + Q_KEM // 2) // Q_KEM) & ((1 << d) - 1) for v in x]
def decompress(d, y):
return [(v * Q_KEM + (1 << (d - 1))) >> d for v in y]
def padd(a, b): return [(x + y) % Q_KEM for x, y in zip(a, b)]
def psub(a, b): return [(x - y) % Q_KEM for x, y in zip(a, b)]
# Round trips, and the compression error bound.
rng = random.Random(21)
for d in (1, 4, 5, 10, 11, 12):
v = [rng.randrange(1 << d) for _ in range(256)]
if d == 12:
v = [x % Q_KEM for x in v]
assert byte_decode(d, byte_encode(d, v)) == v
assert len(byte_encode(d, v)) == 32 * d
print("byte_encode / byte_decode round trip: ok for d in {1,4,5,10,11,12}")
x = [rng.randrange(Q_KEM) for _ in range(256)]
for d in (10, 11, 4, 5):
err = max(abs(modpm(a - b, Q_KEM)) for a, b in zip(x, decompress(d, compress(d, x))))
print(" compress_%-2d worst error %3d (bound q/2^%d = %.1f)"
% (d, err, d + 1, Q_KEM / 2 ** (d + 1)))
# Concrete Number Demonstration: Compression, Bit Reduction, and Reconstruction Error
poly_sample = [rng.randrange(Q_KEM) for _ in range(256)]
comp_d4 = compress(4, poly_sample)
decomp_d4 = decompress(4, comp_d4)
err_d4 = [abs(modpm(orig_val - rec_val, Q_KEM)) for orig_val, rec_val in zip(poly_sample, decomp_d4)]
comp_d10 = compress(10, poly_sample)
decomp_d10 = decompress(10, comp_d10)
err_d10 = [abs(modpm(orig_val - rec_val, Q_KEM)) for orig_val, rec_val in zip(poly_sample, decomp_d10)]
show_card("Polynomial Coefficient Compression & Decompression (Concrete Numbers)",
[("Original Polynomial p[:8] (mod 3329)", chips(poly_sample[:8])),
("Compressed to d=4 bits (values in [0, 15])", chips(comp_d4[:8])),
("Decompressed from d=4 bits", chips(decomp_d4[:8])),
("Compression Errors |p - Decomp(Comp(p))|[:8]", chips(err_d4[:8], highlight=('#f9e2af', 'rgba(223,142,29,0.15)'))),
("Max Error Observed (d=4)", f"{max(err_d4)} levels (Theoretical Bound: ceil(q/2^5) = {math.ceil(Q_KEM / 32)}) " + badge("Small Error Tolerated", "verified")),
("Compressed to d=10 bits (values in [0, 1023])", chips(comp_d10[:8])),
("Decompressed from d=10 bits", chips(decomp_d10[:8])),
("Max Error Observed (d=10)", f"{max(err_d10)} levels (Theoretical Bound: ceil(q/2^11) = {math.ceil(Q_KEM / 2048)}) " + badge("Near-Exact Recovery", "verified"))],
category='token', badge_text="Lossy Compression",
footer="Takeaway: Compression shrinks the ciphertext by discarding lower bits. The resulting rounding difference is just additional noise, which the lattice error-correction margin easily absorbs.")
byte_encode / byte_decode round trip: ok for d in {1,4,5,10,11,12}
compress_10 worst error 2 (bound q/2^11 = 1.6)
compress_11 worst error 1 (bound q/2^12 = 0.8)
compress_4 worst error 104 (bound q/2^5 = 104.0)
compress_5 worst error 52 (bound q/2^6 = 52.0)
PASSED (1.4 ms)
class MLKEM:
"""FIPS 203. Readable, not fast, not constant time."""
def __init__(self, level=768):
p = KEM_PARAMS[level]
self.level, self.k = level, p["k"]
self.eta1, self.eta2 = p["eta1"], p["eta2"]
self.du, self.dv = p["du"], p["dv"]
self.ek_bytes = 384 * self.k + 32
self.dk_bytes = 768 * self.k + 96
self.ct_bytes = 32 * (self.du * self.k + self.dv)
# ---- helpers -------------------------------------------------------
def _expand_a(self, rho):
"""A-hat[i][j] = SampleNTT(rho || j || i). Note the index order."""
return [[sample_ntt(rho + bytes([j, i])) for j in range(self.k)]
for i in range(self.k)]
def _mat_vec(self, mat, vec, transpose=False):
out = []
for i in range(self.k):
acc = [0] * 256
for j in range(self.k):
a = mat[j][i] if transpose else mat[i][j]
acc = padd(acc, pointwise_kem(a, vec[j]))
out.append(acc)
return out
# ---- K-PKE ------------------------------------------------------------
def pke_keygen(self, d):
rho, sigma = kem_G(d + bytes([self.k]))
a = self._expand_a(rho)
n = 0
s = []
for _ in range(self.k):
s.append(sample_poly_cbd(self.eta1, kem_prf(self.eta1, sigma, n))); n += 1
e = []
for _ in range(self.k):
e.append(sample_poly_cbd(self.eta1, kem_prf(self.eta1, sigma, n))); n += 1
s_hat = [ntt_kem(p) for p in s]
e_hat = [ntt_kem(p) for p in e]
t_hat = [padd(x, y) for x, y in zip(self._mat_vec(a, s_hat), e_hat)]
return (b"".join(byte_encode(12, p) for p in t_hat) + rho,
b"".join(byte_encode(12, p) for p in s_hat))
def pke_encrypt(self, ek, m, r):
t_hat = [byte_decode(12, ek[384*i:384*(i+1)]) for i in range(self.k)]
rho = ek[384*self.k:384*self.k+32]
a = self._expand_a(rho)
n = 0
y = []
for _ in range(self.k):
y.append(sample_poly_cbd(self.eta1, kem_prf(self.eta1, r, n))); n += 1
e1 = []
for _ in range(self.k):
e1.append(sample_poly_cbd(self.eta2, kem_prf(self.eta2, r, n))); n += 1
e2 = sample_poly_cbd(self.eta2, kem_prf(self.eta2, r, n))
y_hat = [ntt_kem(p) for p in y]
u = [padd(INTT_kem(p), q) for p, q in
zip(self._mat_vec(a, y_hat, transpose=True), e1)]
mu = decompress(1, byte_decode(1, m))
tv = [0] * 256
for th, yh in zip(t_hat, y_hat):
tv = padd(tv, pointwise_kem(th, yh))
v = padd(padd(INTT_kem(tv), e2), mu)
return (b"".join(byte_encode(self.du, compress(self.du, p)) for p in u)
+ byte_encode(self.dv, compress(self.dv, v)))
def pke_decrypt(self, dk, c):
cut, step = 32 * self.du * self.k, 32 * self.du
u = [decompress(self.du, byte_decode(self.du, c[i*step:(i+1)*step]))
for i in range(self.k)]
v = decompress(self.dv, byte_decode(self.dv, c[cut:]))
s_hat = [byte_decode(12, dk[384*i:384*(i+1)]) for i in range(self.k)]
su = [0] * 256
for sh, ui in zip(s_hat, u):
su = padd(su, pointwise_kem(sh, ntt_kem(ui)))
return byte_encode(1, compress(1, psub(v, INTT_kem(su))))
# ---- ML-KEM -----------------------------------------------------------
def keygen_internal(self, d, z):
ek, dk_pke = self.pke_keygen(d)
return ek, dk_pke + ek + kem_H(ek) + z
def encaps_internal(self, ek, m):
key, r = kem_G(m + kem_H(ek))
return key, self.pke_encrypt(ek, m, r)
def decaps_internal(self, dk, c):
k3 = 384 * self.k
dk_pke = dk[:k3]
ek_pke = dk[k3:768*self.k+32]
h = dk[768*self.k+32:768*self.k+64]
z = dk[768*self.k+64:768*self.k+96]
m2 = self.pke_decrypt(dk_pke, c)
key2, r2 = kem_G(m2 + h)
fallback = kem_J(z + c) # implicit rejection key
return key2 if self.pke_encrypt(ek_pke, m2, r2) == c else fallback
# ---- public API, with the input validation FIPS 203 requires -------
def keygen(self, rng=None):
rng = rng or os.urandom
return self.keygen_internal(rng(32), rng(32))
def encaps(self, ek, rng=None):
rng = rng or os.urandom
if len(ek) != self.ek_bytes:
raise ValueError("ek must be %d bytes, got %d" % (self.ek_bytes, len(ek)))
body = ek[:384 * self.k] # modulus check
if b"".join(byte_encode(12, byte_decode(12, body[384*i:384*(i+1)]))
for i in range(self.k)) != body:
raise ValueError("ek has coefficients outside [0, q)")
return self.encaps_internal(ek, rng(32))
def decaps(self, dk, c):
if len(c) != self.ct_bytes:
raise ValueError("ciphertext must be %d bytes" % self.ct_bytes)
if len(dk) != self.dk_bytes:
raise ValueError("dk must be %d bytes" % self.dk_bytes)
if kem_H(dk[384*self.k:768*self.k+32]) != dk[768*self.k+32:768*self.k+64]:
raise ValueError("dk failed its embedded hash check")
return self.decaps_internal(dk, c)
print("ML-KEM implemented.")
ML-KEM implemented. PASSED (0.8 ms)
# ---- the test suite ------------------------------------------------------
def test_mlkem(verbose=True):
checks = []
def ok(name, cond):
checks.append((name, bool(cond)))
if verbose:
print(" %-52s %s" % (name, "ok" if cond else "FAILED"))
ok("zeta has order 256 mod q", pow(17, 256, Q_KEM) == 1 and pow(17, 128, Q_KEM) == Q_KEM - 1)
rng = random.Random(1)
a = [rng.randrange(Q_KEM) for _ in range(256)]
b = [rng.randrange(Q_KEM) for _ in range(256)]
ok("NTT round trip", INTT_kem(ntt_kem(a)) == a)
ok("NTT multiply == schoolbook",
INTT_kem(pointwise_kem(ntt_kem(a), ntt_kem(b))) == poly_mul_schoolbook(a, b, Q_KEM))
for lvl, (pk, sk, ct) in ((512, (800, 1632, 768)),
(768, (1184, 2400, 1088)),
(1024, (1568, 3168, 1568))):
kem = MLKEM(lvl)
ek, dk = kem.keygen()
key, c = kem.encaps(ek)
ok("ML-KEM-%d sizes match FIPS 203" % lvl,
(len(ek), len(dk), len(c)) == (pk, sk, ct))
ok("ML-KEM-%d shared secrets match" % lvl, kem.decaps(dk, c) == key)
kem = MLKEM(512)
ek, dk = kem.keygen()
key, c = kem.encaps(ek)
bad = bytearray(c); bad[7] ^= 0x01
alt = kem.decaps(dk, bytes(bad))
ok("tampered ciphertext yields a different key", alt != key)
ok("tampered ciphertext still yields 32 bytes", len(alt) == 32)
ok("implicit rejection is deterministic", alt == kem.decaps(dk, bytes(bad)))
d = z = bytes(32); m = bytes(32)
ok("KeyGen is a deterministic function of (d, z)",
kem.keygen_internal(d, z) == kem.keygen_internal(d, z))
ek0, _ = kem.keygen_internal(d, z)
ok("Encaps is a deterministic function of (ek, m)",
kem.encaps_internal(ek0, m) == kem.encaps_internal(ek0, m))
caught = 0
for bad_input in (b"", b"\x00" * (kem.ek_bytes - 1)):
try:
kem.encaps(bad_input)
except ValueError:
caught += 1
body = bytearray(ek0)
body[0:2] = (Q_KEM + 5).to_bytes(2, "little") # a coefficient >= q
try:
kem.encaps(bytes(body))
except ValueError:
caught += 1
ok("input validation rejects bad lengths and out-of-range keys", caught == 3)
fails = 0
kem = MLKEM(768)
for _ in range(25):
ek, dk = kem.keygen()
key, c = kem.encaps(ek)
fails += kem.decaps(dk, c) != key
ok("25 fresh ML-KEM-768 round trips, no failures", fails == 0)
n_ok = sum(1 for _, c in checks if c)
print("\n%d of %d checks passed." % (n_ok, len(checks)))
return n_ok == len(checks)
assert test_mlkem()
zeta has order 256 mod q ok NTT round trip ok NTT multiply == schoolbook ok ML-KEM-512 sizes match FIPS 203 ok ML-KEM-512 shared secrets match ok ML-KEM-768 sizes match FIPS 203 ok ML-KEM-768 shared secrets match ok ML-KEM-1024 sizes match FIPS 203 ok ML-KEM-1024 shared secrets match ok tampered ciphertext yields a different key ok tampered ciphertext still yields 32 bytes ok implicit rejection is deterministic ok KeyGen is a deterministic function of (d, z) ok Encaps is a deterministic function of (ek, m) ok input validation rejects bad lengths and out-of-range keys ok 25 fresh ML-KEM-768 round trips, no failures ok 16 of 16 checks passed. PASSED (236.9 ms)
# ---- the narrative demo: intermediate values (s, e, t, y, u, v) ---------
# We step through ML-KEM-768 with intermediate values displayed so the
# mathematical objects from the slides (s, e, t, y, e1, e2, u, v) are visible.
kem = MLKEM(768)
keccak_reset()
def signed(c, q=3329):
"""Format a coefficient in [-q//2, q//2] for readable inspection."""
return c if c <= q // 2 else c - q
def fmt_poly(p, n=8, as_signed=True):
vals = [signed(c) if as_signed else c for c in p[:n]]
return "[" + ", ".join(f"{x:+d}" if as_signed else str(x) for x in vals) + f", ... (total {len(p)} coeffs)]"
print("=" * 78)
print("1. ALICE KEY GENERATION (FIPS 203 Section 5.1 / K-PKE.KeyGen)")
print("=" * 78)
d = os.urandom(32)
z = os.urandom(32)
ek, dk = kem.keygen_internal(d, z)
rho, sigma = kem_G(d + bytes([kem.k]))
a_hat = kem._expand_a(rho)
# Sample private small secrets s and errors e from CBD(eta1=2)
s, e = [], []
ctr = 0
for _ in range(kem.k):
s.append(sample_poly_cbd(kem.eta1, kem_prf(kem.eta1, sigma, ctr))); ctr += 1
for _ in range(kem.k):
e.append(sample_poly_cbd(kem.eta1, kem_prf(kem.eta1, sigma, ctr))); ctr += 1
s_hat = [ntt_kem(p) for p in s]
t_hat = [byte_decode(12, ek[384*i:384*(i+1)]) for i in range(kem.k)]
print(f"[*] Seed rho (32 bytes): {rho.hex()[:24]}...")
print(f" Matrix A-hat shape: {len(a_hat)} rows x {len(a_hat[0])} cols of polynomials in NTT domain (each row: {len(a_hat[0])} x 256 coeffs in R_q)")
print(f"[*] Secret vector s in R_q^{kem.k}:")
print(f" Shape: column vector of {len(s)} polynomials, each of {len(s[0])} small coefficients in [-eta1, eta1] = [-2, 2]")
for i in range(kem.k):
print(f" s[{i}][:8] = {fmt_poly(s[i])} range: [{min(signed(x) for x in s[i])}, {max(signed(x) for x in s[i])}]")
print(f"[*] Error vector e in R_q^{kem.k}:")
print(f" Shape: column vector of {len(e)} polynomials, each of {len(e[0])} small coefficients in [-eta1, eta1] = [-2, 2]")
for i in range(kem.k):
print(f" e[{i}][:8] = {fmt_poly(e[i])} range: [{min(signed(x) for x in e[i])}, {max(signed(x) for x in e[i])}]")
print(f"[*] Public key vector t-hat = A-hat o s-hat + e-hat mod q:")
print(f" Shape: column vector of {len(t_hat)} polynomials in NTT domain, each of {len(t_hat[0])} coefficients in [0, 3329)")
for i in range(kem.k):
print(f" t_hat[{i}][:8] = {fmt_poly(t_hat[i], as_signed=False)}")
print(f"[*] Published ek: {len(ek)} bytes (t-hat: {384*kem.k} B + rho: 32 B) | Alice dk: {len(dk)} bytes\n")
print("=" * 78)
print("2. BOB ENCAPSULATION / ENCRYPTION (FIPS 203 Section 5.2 / K-PKE.Encrypt)")
print("=" * 78)
m = os.urandom(32)
key_bob, r = kem_G(m + kem_H(ek))
# Ephemeral secret y and errors e1, e2
y, e1 = [], []
ctr = 0
for _ in range(kem.k):
y.append(sample_poly_cbd(kem.eta1, kem_prf(kem.eta1, r, ctr))); ctr += 1
for _ in range(kem.k):
e1.append(sample_poly_cbd(kem.eta2, kem_prf(kem.eta2, r, ctr))); ctr += 1
e2 = sample_poly_cbd(kem.eta2, kem_prf(kem.eta2, r, ctr))
y_hat = [ntt_kem(p) for p in y]
# Message embedding mu in R_q: 0 -> 0, 1 -> round(q/2) = 1665
mu = decompress(1, byte_decode(1, m))
# Ciphertext components: u = A^T y + e1, v = t^T y + e2 + mu
u = [padd(INTT_kem(p), q_poly) for p, q_poly in
zip(kem._mat_vec(a_hat, y_hat, transpose=True), e1)]
tv = [0] * 256
for th, yh in zip(t_hat, y_hat):
tv = padd(tv, pointwise_kem(th, yh))
v = padd(padd(INTT_kem(tv), e2), mu)
ct = kem.pke_encrypt(ek, m, r)
print(f"[*] Plaintext message m: 32 bytes ({32*8} bits)")
print(f"[*] Encoded message mu in R_q:")
print(f" Shape: 1 scalar polynomial of {len(mu)} coefficients (bits scaled to 0 or round(q/2)=1665)")
print(f" mu[:8] = {fmt_poly(mu, as_signed=False)}")
print(f"[*] Ephemeral secret vector y in R_q^{kem.k}:")
print(f" Shape: column vector of {len(y)} polynomials, each of {len(y[0])} small coefficients in [-eta1, eta1] = [-2, 2]")
for i in range(kem.k):
print(f" y[{i}][:8] = {fmt_poly(y[i])}")
print(f"[*] Ciphertext vector u = A^T y + e1 mod q:")
print(f" Shape: column vector of {len(u)} polynomials, each of {len(u[0])} coefficients in [0, 3329)")
for i in range(kem.k):
print(f" u[{i}][:8] = {fmt_poly(u[i], as_signed=False)}")
print(f"[*] Ciphertext scalar polynomial v = t^T y + e2 + mu mod q:")
print(f" Shape: 1 scalar polynomial of {len(v)} coefficients in [0, 3329)")
print(f" v[:8] = {fmt_poly(v, as_signed=False)}")
print(f"[*] Compressed wire ciphertext c = (Compress_10(u), Compress_4(v)): {len(ct)} bytes")
print(f" u compressed: {kem.k} polynomials x 256 coeffs x 10 bits = {kem.k * 320} bytes")
print(f" v compressed: 1 polynomial x 256 coeffs x 4 bits = 128 bytes (total {len(ct)} bytes)")
# ---- [ITEM 1] LOSSY COMPRESSION ERROR INSPECTION (Delta c) ---------------
u0_raw = u[0][0]
u0_comp = compress(kem.du, [u0_raw])[0]
u0_decomp = decompress(kem.du, [u0_comp])[0]
u0_err = signed(u0_raw - u0_decomp)
v0_raw = v[0]
v0_comp = compress(kem.dv, [v0_raw])[0]
v0_decomp = decompress(kem.dv, [v0_comp])[0]
v0_err = signed(v0_raw - v0_decomp)
print(f"[*] [Item 1] Lossy Compression Error (Delta c):")
print(f" u[0][0]: raw={u0_raw:4d} (12-bit) -> comp={u0_comp:3d} (10-bit) -> decomp={u0_decomp:4d} | error Delta u = {u0_err:+d} (bound <= 2)")
print(f" v[0]: raw={v0_raw:4d} (12-bit) -> comp={v0_comp:2d} ( 4-bit) -> decomp={v0_decomp:4d} | error Delta v = {v0_err:+d} (bound <= 104)")
print(f"[*] Bob shared secret: {key_bob.hex()}\n")
print("=" * 78)
print("3. ALICE DECAPSULATION / DECRYPTION (FIPS 203 Section 5.3 / K-PKE.Decrypt)")
print("=" * 78)
# Decompress wire ciphertext
cut, step = 32 * kem.du * kem.k, 32 * kem.du
u_rec = [decompress(kem.du, byte_decode(kem.du, ct[i*step:(i+1)*step])) for i in range(kem.k)]
v_rec = decompress(kem.dv, byte_decode(kem.dv, ct[cut:]))
# Compute v - s^T u = mu + noise
su = [0] * 256
for sh, ui in zip(s_hat, u_rec):
su = padd(su, pointwise_kem(sh, ntt_kem(ui)))
diff = psub(v_rec, INTT_kem(su))
# Effective noise: distance from nearest message point (0 or 1665)
all_noise = [signed((d - target) % 3329) for d, target in zip(diff, mu)]
m_rec = byte_encode(1, compress(1, diff))
key_alice = kem.decaps(dk, ct)
print(f"[*] Recomputed noisy combination: diff = v - s^T u (shape: 1 polynomial of {len(diff)} coeffs in R_q):")
print(f" diff[:8] = {fmt_poly(diff, as_signed=False)}")
print(f" target mu[:8]= {fmt_poly(mu, as_signed=False)}")
print(f" noise[:8] = {all_noise[:8]}")
# ---- [ITEM 2] NOISE BUDGET BREAKDOWN -------------------------------------
max_noise = max(abs(x) for x in all_noise)
mean_noise = sum(abs(x) for x in all_noise) / len(all_noise)
q4_bound = 3329 // 4 # 832
print(f"[*] [Item 2] Decryption Noise Budget Breakdown:")
print(f" Theoretical error formula: e_total = (e2 + e^T y - s^T e1) + (Delta v - s^T Delta u)")
print(f" Observed max |noise|: {max_noise} (mean: {mean_noise:.1f})")
print(f" Decoding failure threshold: q / 4 = {q4_bound}")
print(f" Safety headroom remaining: {q4_bound - max_noise} (failure probability: 2^-164.8, effectively 0)")
print(f"[*] 1-bit rounded message bits: match original m? {m_rec == m}")
print(f"[*] Alice decapsulated key: {key_alice.hex()}")
print(f"[*] Shared secret match: {key_alice == key_bob}\n")
print(f"Total on the wire for this handshake: {len(ek) + len(ct)} bytes")
print(f"The same handshake with X25519: {32 + 32} bytes")
print(f"Keccak: {KECCAK['calls']} calls, about {KECCAK['perms']} permutations\n")
# Rich Summary Cards for ML-KEM-768
show_card("ML-KEM-768 Key Generation Summary",
[("Seed rho (32 bytes)", f"{rho.hex()[:16]}..."),
("Secret Vector s[0][:8]", chips([signed(x) for x in s[0][:8]])),
("Noise Vector e[0][:8]", chips([signed(x) for x in e[0][:8]])),
("Public Key t-hat[0][:8]", chips(t_hat[0][:8])),
("Published ek Size", f"{len(ek)} bytes (t-hat: {384*kem.k} B + rho: 32 B)")],
category='secret', badge_text="KeyGen Stage")
show_card("ML-KEM-768 Encapsulation Summary",
[("Plaintext Message m", f"{m.hex()[:16]}..."),
("Ephemeral Vector y[0][:8]", chips([signed(x) for x in y[0][:8]])),
("Ciphertext Vector u[0][:8]", chips(u[0][:8])),
("Ciphertext Scalar v[:8]", chips(v[:8])),
("Compressed Ciphertext Size", f"{len(ct)} bytes (u: {kem.k*320} B, v: 128 B)"),
("Bob Shared Secret K_bob", f"{key_bob.hex()[:24]}...")],
category='token', badge_text="Encaps Stage")
show_card("ML-KEM-768 Decapsulation & Agreement",
[("Observed Max Noise", f"{max_noise} (Bound q/4 = {q4_bound})"),
("Safety Headroom Remaining", f"{q4_bound - max_noise} levels"),
("Decoded Message Match", f"m_rec == m: {m_rec == m} " + badge("Message Recovered", "verified")),
("Alice Shared Secret K_alice", f"{key_alice.hex()[:24]}..."),
("Key Synchronization Match", f"K_alice == K_bob: {key_alice == key_bob} " + badge("Keys Synchronized", "verified"))],
category='verified', badge_text="Decaps Stage")
==============================================================================
1. ALICE KEY GENERATION (FIPS 203 Section 5.1 / K-PKE.KeyGen)
==============================================================================
[*] Seed rho (32 bytes): a75da6ea6bd3360b6d8ce4e7...
Matrix A-hat shape: 3 rows x 3 cols of polynomials in NTT domain (each row: 3 x 256 coeffs in R_q)
[*] Secret vector s in R_q^3:
Shape: column vector of 3 polynomials, each of 256 small coefficients in [-eta1, eta1] = [-2, 2]
s[0][:8] = [+0, +0, -1, -1, -1, -1, +2, -1, ... (total 256 coeffs)] range: [-2, 2]
s[1][:8] = [-1, +0, +1, -1, -1, +0, -2, +0, ... (total 256 coeffs)] range: [-2, 2]
s[2][:8] = [+1, -2, +0, +1, -1, +1, +0, -1, ... (total 256 coeffs)] range: [-2, 2]
[*] Error vector e in R_q^3:
Shape: column vector of 3 polynomials, each of 256 small coefficients in [-eta1, eta1] = [-2, 2]
e[0][:8] = [+1, -1, +1, +1, +0, +1, -1, +0, ... (total 256 coeffs)] range: [-2, 2]
e[1][:8] = [+1, -1, +1, -2, +0, -2, +0, +0, ... (total 256 coeffs)] range: [-2, 2]
e[2][:8] = [-1, +1, +1, +2, -1, -1, -1, -2, ... (total 256 coeffs)] range: [-2, 2]
[*] Public key vector t-hat = A-hat o s-hat + e-hat mod q:
Shape: column vector of 3 polynomials in NTT domain, each of 256 coefficients in [0, 3329)
t_hat[0][:8] = [2260, 1352, 3106, 792, 694, 2253, 2095, 128, ... (total 256 coeffs)]
t_hat[1][:8] = [1960, 3145, 128, 1760, 1785, 791, 2864, 2024, ... (total 256 coeffs)]
t_hat[2][:8] = [1814, 1713, 1825, 606, 1655, 808, 3222, 1186, ... (total 256 coeffs)]
[*] Published ek: 1184 bytes (t-hat: 1152 B + rho: 32 B) | Alice dk: 2400 bytes
==============================================================================
2. BOB ENCAPSULATION / ENCRYPTION (FIPS 203 Section 5.2 / K-PKE.Encrypt)
==============================================================================
[*] Plaintext message m: 32 bytes (256 bits)
[*] Encoded message mu in R_q:
Shape: 1 scalar polynomial of 256 coefficients (bits scaled to 0 or round(q/2)=1665)
mu[:8] = [0, 0, 0, 1665, 1665, 1665, 1665, 1665, ... (total 256 coeffs)]
[*] Ephemeral secret vector y in R_q^3:
Shape: column vector of 3 polynomials, each of 256 small coefficients in [-eta1, eta1] = [-2, 2]
y[0][:8] = [-1, -1, +0, -1, +0, +2, +0, +1, ... (total 256 coeffs)]
y[1][:8] = [-1, +1, -1, +1, +1, +0, +0, +1, ... (total 256 coeffs)]
y[2][:8] = [+0, +0, +1, +0, -1, +0, +0, +2, ... (total 256 coeffs)]
[*] Ciphertext vector u = A^T y + e1 mod q:
Shape: column vector of 3 polynomials, each of 256 coefficients in [0, 3329)
u[0][:8] = [1321, 1327, 250, 393, 3327, 3033, 1754, 2216, ... (total 256 coeffs)]
u[1][:8] = [1085, 1055, 2292, 1438, 2842, 3243, 2120, 543, ... (total 256 coeffs)]
u[2][:8] = [948, 520, 2654, 1088, 1834, 605, 288, 2832, ... (total 256 coeffs)]
[*] Ciphertext scalar polynomial v = t^T y + e2 + mu mod q:
Shape: 1 scalar polynomial of 256 coefficients in [0, 3329)
v[:8] = [985, 1974, 2346, 1322, 246, 1678, 574, 1121, ... (total 256 coeffs)]
[*] Compressed wire ciphertext c = (Compress_10(u), Compress_4(v)): 1088 bytes
u compressed: 3 polynomials x 256 coeffs x 10 bits = 960 bytes
v compressed: 1 polynomial x 256 coeffs x 4 bits = 128 bytes (total 1088 bytes)
[*] [Item 1] Lossy Compression Error (Delta c):
u[0][0]: raw=1321 (12-bit) -> comp=406 (10-bit) -> decomp=1320 | error Delta u = +1 (bound <= 2)
v[0]: raw= 985 (12-bit) -> comp= 5 ( 4-bit) -> decomp=1040 | error Delta v = -55 (bound <= 104)
[*] Bob shared secret: caf076540ad507f454f3d33863a2fea3de6256b7ff76af114100ef4e7301c4af
==============================================================================
3. ALICE DECAPSULATION / DECRYPTION (FIPS 203 Section 5.3 / K-PKE.Decrypt)
==============================================================================
[*] Recomputed noisy combination: diff = v - s^T u (shape: 1 polynomial of 256 coeffs in R_q):
diff[:8] = [52, 3312, 3302, 1580, 1556, 1731, 1633, 1623, ... (total 256 coeffs)]
target mu[:8]= [0, 0, 0, 1665, 1665, 1665, 1665, 1665, ... (total 256 coeffs)]
noise[:8] = [52, -17, -27, -85, -109, 66, -32, -42]
[*] [Item 2] Decryption Noise Budget Breakdown:
Theoretical error formula: e_total = (e2 + e^T y - s^T e1) + (Delta v - s^T Delta u)
Observed max |noise|: 211 (mean: 61.1)
Decoding failure threshold: q / 4 = 832
Safety headroom remaining: 621 (failure probability: 2^-164.8, effectively 0)
[*] 1-bit rounded message bits: match original m? True
[*] Alice decapsulated key: caf076540ad507f454f3d33863a2fea3de6256b7ff76af114100ef4e7301c4af
[*] Shared secret match: True
Total on the wire for this handshake: 2272 bytes
The same handshake with X25519: 64 bytes
Keccak: 77 calls, about 181 permutations
PASSED (13.1 ms)
# ---- WHAT BREAKS ML-KEM: ERROR INJECTION & DECAPSULATION FAILURE ----------
kem = MLKEM(768)
ek, dk = kem.keygen()
m_orig = os.urandom(32)
key_bob_true, ct_clean = kem.encaps_internal(ek, m_orig)
cut = 32 * kem.du * kem.k
u_bytes = ct_clean[:cut]
v_bytes = ct_clean[cut:]
v_clean = decompress(kem.dv, byte_decode(kem.dv, v_bytes))
v_faulted = list(v_clean)
v_faulted[0] = (v_faulted[0] + 1665) % 3329
ct_faulted = u_bytes + byte_encode(kem.dv, compress(kem.dv, v_faulted))
m_dec_clean = kem.pke_decrypt(dk[:384*kem.k], ct_clean)
m_dec_faulted = kem.pke_decrypt(dk[:384*kem.k], ct_faulted)
key_alice_faulted = kem.decaps(dk, ct_faulted)
show_card("Noise Injection Attack: Perturbing Ciphertext (+1665 mod 3329)",
[("Original Message Bit 0", chip(m_orig[0] & 1)),
("Decoded Bit 0 (clean ct)", chip(m_dec_clean[0] & 1) + " " + badge("Correct", "verified")),
("Decoded Bit 0 (faulted ct)", chip(m_dec_faulted[0] & 1) + " " + badge("Bit Flipped", "reject")),
("Bob True Shared Key", f"{key_bob_true.hex()[:24]}..."),
("Alice Decapsulated Key", f"{key_alice_faulted.hex()[:24]}..."),
("Key Synchronization Result", f"Match? {key_alice_faulted == key_bob_true} " + badge("Silent Desynchronization (FO Trapdoor)", "reject"))],
category='reject', badge_text="FO Defense")
noise_rows = []
for delta in (0, 200, 400, 600, 750, 850, 950, 1665):
fail_pke, fail_kem = 0, 0
trials = 60
for _ in range(trials):
ek_, dk_ = kem.keygen()
m_ = os.urandom(32)
key_b, ct_ = kem.encaps_internal(ek_, m_)
v_dec = decompress(kem.dv, byte_decode(kem.dv, ct_[cut:]))
v_dec[0] = (v_dec[0] + delta) % 3329
ct_pert = ct_[:cut] + byte_encode(kem.dv, compress(kem.dv, v_dec))
m_pert = kem.pke_decrypt(dk_[:384*kem.k], ct_pert)
if m_pert != m_: fail_pke += 1
key_a = kem.decaps(dk_, ct_pert)
if key_a != key_b: fail_kem += 1
if delta == 0:
st = badge("Honest (Synced)", "verified")
elif fail_pke == 0:
st = badge("Tamper Caught by FO", "input")
else:
st = badge(f"Noise Overflow ({100*fail_pke/trials:.0f}% Flip)", "reject")
noise_rows.append([f"+{delta} mod q", f"{fail_pke/trials*100:5.1f} %", f"{fail_kem/trials*100:5.1f} %", st])
show_table(["Injected Error (delta)", "Message Bit Flip Rate", "FO Implicit Rejection", "Observed Security Behavior"],
noise_rows, title="ML-KEM Noise Sweep & Fujisaki-Okamoto Defense (60 trials each)", category='reject')
PASSED (3976.9 ms)
| Injected Error (delta) | Message Bit Flip Rate | FO Implicit Rejection | Observed Security Behavior |
|---|---|---|---|
| +0 mod q | 0.0 % | 0.0 % | Honest (Synced) |
| +200 mod q | 0.0 % | 100.0 % | Tamper Caught by FO |
| +400 mod q | 0.0 % | 100.0 % | Tamper Caught by FO |
| +600 mod q | 0.0 % | 100.0 % | Tamper Caught by FO |
| +750 mod q | 40.0 % | 100.0 % | Noise Overflow (40% Flip) |
| +850 mod q | 46.7 % | 100.0 % | Noise Overflow (47% Flip) |
| +950 mod q | 100.0 % | 100.0 % | Noise Overflow (100% Flip) |
| +1665 mod q | 100.0 % | 100.0 % | Noise Overflow (100% Flip) |
# ---- implicit rejection, the part people get wrong ----------------------
print("Flip one bit of the ciphertext and decapsulate again.\n")
for pos in (0, 100, 500, len(ct) - 1):
bad = bytearray(ct); bad[pos] ^= 0x01
k = kem.decaps(dk, bytes(bad))
print(" bit flipped at byte %4d -> %s %s"
% (pos, k.hex()[:32] + "...", "MATCHES (bad!)" if k == key_alice else "different key"))
print()
print("Notice what did *not* happen: no exception, no error code, no boolean.")
print("Decaps returned a perfectly ordinary-looking 32-byte key derived from a")
print("secret z inside dk. The handshake will fail later, when the two sides")
print("find their transcripts do not authenticate, and nothing leaked.")
print()
print("An implementation that returns an error here hands the attacker a")
print("chosen-ciphertext oracle, and the private key follows.")
Flip one bit of the ciphertext and decapsulate again. bit flipped at byte 0 -> 50d26fb918168776cc7d51fbc6b9cc23... different key bit flipped at byte 100 -> 2adc838cd00421c891782c39bb23a3c0... different key bit flipped at byte 500 -> 77fb476d4166162ccfc9fff2d7a07a7b... different key bit flipped at byte 1087 -> f5a2620ebeae9b295a460e5ccc526087... different key Notice what did *not* happen: no exception, no error code, no boolean. Decaps returned a perfectly ordinary-looking 32-byte key derived from a secret z inside dk. The handshake will fail later, when the two sides find their transcripts do not authenticate, and nothing leaked. An implementation that returns an error here hands the attacker a chosen-ciphertext oracle, and the private key follows. PASSED (13.2 ms)
# ---- a hook for the official test vectors -------------------------------
def run_acvp_mlkem(path):
"""Check against NIST ACVP vectors, if you have downloaded them.
Get ML-KEM keyGen / encapDecap prompt+expected JSON from
https://github.com/usnistgov/ACVP-Server (gen-val/json-files) and point
this at the directory. Nothing here needs the network.
"""
import glob, json as _json
files = sorted(glob.glob(os.path.join(path, "*.json")))
if not files:
print("No JSON files in %r. Skipping." % path)
return
print("Found %d files. Wire up the group parsing for the ones you care about:" % len(files))
for f in files[:10]:
print(" ", os.path.basename(f))
print("The suite above is self-validating: it checks the NTT against schoolbook,")
print("every codec round trip, all three parameter sets against the published")
print("sizes, and implicit rejection. Those catch essentially any real bug.")
print()
print("For formal conformance, drop the ACVP vectors in and call run_acvp_mlkem().")
The suite above is self-validating: it checks the NTT against schoolbook, every codec round trip, all three parameter sets against the published sizes, and implicit rejection. Those catch essentially any real bug. For formal conformance, drop the ACVP vectors in and call run_acvp_mlkem(). PASSED (0.1 ms)
---
Complete, standard-compliant implementation of ML-DSA (Module-LWE Digital Signature Algorithm).
The core mechanism is Fiat-Shamir with aborts: signing is a rejection loop that retries until the signature values satisfy strict norm bounds, ensuring that no secret key bits ever leak.
# ---- parameters and the rounding functions ------------------------------
D_DROP = 13
DSA_PARAMS = {
44: dict(tau=39, lam=128, gamma1=1 << 17, gamma2=(Q_DSA - 1) // 88,
k=4, l=4, eta=2, omega=80),
65: dict(tau=49, lam=192, gamma1=1 << 19, gamma2=(Q_DSA - 1) // 32,
k=6, l=5, eta=4, omega=55),
87: dict(tau=60, lam=256, gamma1=1 << 19, gamma2=(Q_DSA - 1) // 32,
k=8, l=7, eta=2, omega=75),
}
def dsa_H(data, out_len):
return _shake256(data, out_len)
def bitlen(a):
return a.bit_length()
def inf_norm(poly):
return max(abs(modpm(c, Q_DSA)) for c in poly)
def vec_norm(vec):
return max((inf_norm(p) for p in vec), default=0)
def power2round(r):
"""Split off the top bits of r. Used to shrink the public key."""
rp = r % Q_DSA
r0 = modpm(rp, 1 << D_DROP)
return (rp - r0) >> D_DROP, r0
def decompose(r, gamma2):
"""Split r into a coarse bucket and a signed remainder."""
rp = r % Q_DSA
r0 = modpm(rp, 2 * gamma2)
if rp - r0 == Q_DSA - 1: # the edge case everybody misses
return 0, r0 - 1
return (rp - r0) // (2 * gamma2), r0
def high_bits(r, g2): return decompose(r, g2)[0]
def low_bits(r, g2): return decompose(r, g2)[1]
def make_hint(z, r, g2):
return int(high_bits(r, g2) != high_bits(r + z, g2))
def use_hint(h, r, g2):
m = (Q_DSA - 1) // (2 * g2)
r1, r0 = decompose(r, g2)
if h == 1:
return (r1 + 1) % m if r0 > 0 else (r1 - 1) % m
return r1
# The two identities the whole scheme leans on.
rng = random.Random(5)
for _ in range(3000):
r = rng.randrange(Q_DSA)
h, l = power2round(r)
assert (h * (1 << D_DROP) + l) % Q_DSA == r
for g2 in ((Q_DSA - 1) // 88, (Q_DSA - 1) // 32):
r1, r0 = decompose(r, g2)
assert (r1 * 2 * g2 + r0) % Q_DSA == r
z = rng.randrange(-g2, g2 + 1) % Q_DSA
assert use_hint(make_hint(z, r, g2), r, g2) == high_bits((r + z) % Q_DSA, g2)
print("Power2Round, Decompose and the MakeHint/UseHint identity: verified on 3000 values.")
Power2Round, Decompose and the MakeHint/UseHint identity: verified on 3000 values. PASSED (8.2 ms)
# ---- bit packing --------------------------------------------------------
def simple_bit_pack(w, b):
bits = bitlen(b); acc = 0
for i, v in enumerate(w):
acc |= (v & ((1 << bits) - 1)) << (bits * i)
return acc.to_bytes(32 * bits, "little")
def simple_bit_unpack(v, b):
bits = bitlen(b); acc = int.from_bytes(v, "little"); mask = (1 << bits) - 1
return [(acc >> (bits * i)) & mask for i in range(256)]
def bit_pack(w, a, b):
bits = bitlen(a + b); acc = 0
for i, v in enumerate(w):
acc |= ((b - modpm(v, Q_DSA)) & ((1 << bits) - 1)) << (bits * i)
return acc.to_bytes(32 * bits, "little")
def bit_unpack(v, a, b):
bits = bitlen(a + b); acc = int.from_bytes(v, "little"); mask = (1 << bits) - 1
return [(b - ((acc >> (bits * i)) & mask)) % Q_DSA for i in range(256)]
# ---- sampling -----------------------------------------------------------
def sample_in_ball(rho, tau):
"""Exactly tau coefficients of +/-1, everything else zero."""
c = [0] * 256
st = Squeeze("SHAKE256", rho)
signs = int.from_bytes(st.take(8), "little")
for i in range(256 - tau, 256):
while True:
j = st.take(1)[0]
if j <= i:
break
c[i] = c[j]
c[j] = (Q_DSA - 1) if (signs >> (i - (256 - tau))) & 1 else 1
return c
def rej_ntt_poly(rho):
"""Uniform mod q in the NTT domain, from 23-bit samples."""
st = Squeeze("SHAKE128", rho)
out = []
while len(out) < 256:
b = st.take(3)
z = b[0] + (b[1] << 8) + ((b[2] & 0x7F) << 16)
if z < Q_DSA:
out.append(z)
return out
def rej_bounded_poly(rho, eta):
"""Coefficients in [-eta, eta], from half-bytes."""
st = Squeeze("SHAKE256", rho)
out = []
while len(out) < 256:
z = st.take(1)[0]
for half in (z & 0x0F, z >> 4):
if len(out) == 256:
break
if eta == 2 and half < 15:
out.append((2 - (half % 5)) % Q_DSA)
elif eta == 4 and half < 9:
out.append((4 - half) % Q_DSA)
return out
print("ML-DSA rounding, packing and samplers ready.")
ML-DSA rounding, packing and samplers ready. PASSED (0.3 ms)
class MLDSA:
"""FIPS 204. Readable, not fast, not constant time."""
def __init__(self, level=65):
for name, value in DSA_PARAMS[level].items():
setattr(self, name, value)
self.level = level
self.beta = self.tau * self.eta
self.c_tilde_bytes = self.lam // 4
self.z_bits = bitlen(2 * self.gamma1 - 1)
self.t1_bits = bitlen(Q_DSA - 1) - D_DROP
self.eta_bits = bitlen(2 * self.eta)
self.pk_bytes = 32 + 32 * self.t1_bits * self.k
self.sk_bytes = 128 + 32 * self.eta_bits * (self.k + self.l) + 32 * D_DROP * self.k
self.sig_bytes = self.c_tilde_bytes + 32 * self.z_bits * self.l + self.omega + self.k
# ---- expansion --------------------------------------------------------
def expand_a(self, rho):
return [[rej_ntt_poly(rho + bytes([s, r])) for s in range(self.l)]
for r in range(self.k)]
def expand_s(self, rho):
s1 = [rej_bounded_poly(rho + r.to_bytes(2, "little"), self.eta)
for r in range(self.l)]
s2 = [rej_bounded_poly(rho + (r + self.l).to_bytes(2, "little"), self.eta)
for r in range(self.k)]
return s1, s2
def expand_mask(self, rho, kappa):
c = 1 + bitlen(self.gamma1 - 1)
return [bit_unpack(dsa_H(rho + (kappa + r).to_bytes(2, "little"), 32 * c),
self.gamma1 - 1, self.gamma1) for r in range(self.l)]
# ---- encoding ---------------------------------------------------------
def pk_encode(self, rho, t1):
return rho + b"".join(simple_bit_pack(p, (1 << self.t1_bits) - 1) for p in t1)
def pk_decode(self, pk):
step = 32 * self.t1_bits
return pk[:32], [simple_bit_unpack(pk[32+i*step:32+(i+1)*step],
(1 << self.t1_bits) - 1)
for i in range(self.k)]
def sk_encode(self, rho, key, tr, s1, s2, t0):
out = rho + key + tr
for p in s1 + s2:
out += bit_pack(p, self.eta, self.eta)
for p in t0:
out += bit_pack(p, (1 << (D_DROP-1)) - 1, 1 << (D_DROP-1))
return out
def sk_decode(self, sk):
rho, key, tr = sk[:32], sk[32:64], sk[64:128]
pos, step = 128, 32 * self.eta_bits
s1, s2 = [], []
for target, count in ((s1, self.l), (s2, self.k)):
for _ in range(count):
target.append(bit_unpack(sk[pos:pos+step], self.eta, self.eta))
pos += step
step, t0 = 32 * D_DROP, []
for _ in range(self.k):
t0.append(bit_unpack(sk[pos:pos+step],
(1 << (D_DROP-1)) - 1, 1 << (D_DROP-1)))
pos += step
return rho, key, tr, s1, s2, t0
def w1_encode(self, w1):
b = (Q_DSA - 1) // (2 * self.gamma2) - 1
return b"".join(simple_bit_pack(p, b) for p in w1)
def hint_pack(self, h):
y = bytearray(self.omega + self.k); idx = 0
for i in range(self.k):
for j in range(256):
if h[i][j]:
y[idx] = j; idx += 1
y[self.omega + i] = idx
return bytes(y)
def hint_unpack(self, y):
h = [[0] * 256 for _ in range(self.k)]; idx = 0
for i in range(self.k):
end = y[self.omega + i]
if end < idx or end > self.omega:
return None
first = idx
while idx < end:
if idx > first and y[idx - 1] >= y[idx]:
return None
h[i][y[idx]] = 1; idx += 1
return None if any(y[j] for j in range(idx, self.omega)) else h
def sig_encode(self, c_tilde, z, h):
return (c_tilde + b"".join(bit_pack(p, self.gamma1 - 1, self.gamma1)
for p in z) + self.hint_pack(h))
def sig_decode(self, sig):
pos, step = self.c_tilde_bytes, 32 * self.z_bits
z = []
for _ in range(self.l):
z.append(bit_unpack(sig[pos:pos+step], self.gamma1 - 1, self.gamma1))
pos += step
return sig[:self.c_tilde_bytes], z, self.hint_unpack(sig[pos:])
# ---- key generation ---------------------------------------------------
def keygen_internal(self, xi):
seed = dsa_H(xi + bytes([self.k, self.l]), 128)
rho, rho_p, key = seed[:32], seed[32:96], seed[96:128]
a_hat = self.expand_a(rho)
s1, s2 = self.expand_s(rho_p)
s1_hat = [ntt_dsa(p) for p in s1]
t1, t0 = [], []
for i in range(self.k):
acc = [0] * 256
for j in range(self.l):
acc = [(x + y) % Q_DSA for x, y in
zip(acc, pointwise_dsa(a_hat[i][j], s1_hat[j]))]
t = [(x + y) % Q_DSA for x, y in zip(INTT_dsa(acc), s2[i])]
hi, lo = zip(*(power2round(c) for c in t))
t1.append(list(hi)); t0.append([x % Q_DSA for x in lo])
pk = self.pk_encode(rho, t1)
return pk, self.sk_encode(rho, key, dsa_H(pk, 64), s1, s2, t0)
def keygen(self, rng=None):
rng = rng or os.urandom
return self.keygen_internal(rng(32))
# ---- signing ----------------------------------------------------------
def sign_internal(self, sk, m_prime, rnd, stats=False):
rho, key, tr, s1, s2, t0 = self.sk_decode(sk)
s1_hat = [ntt_dsa(p) for p in s1]
s2_hat = [ntt_dsa(p) for p in s2]
t0_hat = [ntt_dsa(p) for p in t0]
a_hat = self.expand_a(rho)
mu = dsa_H(tr + m_prime, 64)
rho_pp = dsa_H(key + rnd + mu, 64)
kappa, attempts, why = 0, 0, []
while True:
attempts += 1
y = self.expand_mask(rho_pp, kappa)
kappa += self.l
y_hat = [ntt_dsa(p) for p in y]
w = []
for i in range(self.k):
acc = [0] * 256
for j in range(self.l):
acc = [(x + yy) % Q_DSA for x, yy in
zip(acc, pointwise_dsa(a_hat[i][j], y_hat[j]))]
w.append(INTT_dsa(acc))
w1 = [[high_bits(c, self.gamma2) for c in p] for p in w]
c_tilde = dsa_H(mu + self.w1_encode(w1), self.c_tilde_bytes)
c_hat = ntt_dsa(sample_in_ball(c_tilde, self.tau))
cs1 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in s1_hat]
cs2 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in s2_hat]
z = [[(a + b) % Q_DSA for a, b in zip(p, q)] for p, q in zip(y, cs1)]
wcs2 = [[(a - b) % Q_DSA for a, b in zip(p, q)] for p, q in zip(w, cs2)]
r0 = [[low_bits(c, self.gamma2) for c in p] for p in wcs2]
if vec_norm(z) >= self.gamma1 - self.beta:
why.append("z too large"); continue
if max(max(abs(c) for c in p) for p in r0) >= self.gamma2 - self.beta:
why.append("r0 too large"); continue
ct0 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in t0_hat]
if vec_norm(ct0) >= self.gamma2:
why.append("c*t0 too large"); continue
h = [[make_hint((-ct0[i][j]) % Q_DSA,
(wcs2[i][j] + ct0[i][j]) % Q_DSA, self.gamma2)
for j in range(256)] for i in range(self.k)]
if sum(sum(row) for row in h) > self.omega:
why.append("too many hints"); continue
sig = self.sig_encode(c_tilde, z, h)
return (sig, attempts, why) if stats else sig
def sign(self, sk, message, ctx=b"", deterministic=False, rng=None, stats=False):
if len(ctx) > 255:
raise ValueError("context must be at most 255 bytes")
rng = rng or os.urandom
rnd = bytes(32) if deterministic else rng(32)
return self.sign_internal(sk, bytes([0, len(ctx)]) + ctx + message, rnd, stats)
# ---- verification -----------------------------------------------------
def verify_internal(self, pk, m_prime, sig):
if len(sig) != self.sig_bytes or len(pk) != self.pk_bytes:
return False
rho, t1 = self.pk_decode(pk)
c_tilde, z, h = self.sig_decode(sig)
if h is None or vec_norm(z) >= self.gamma1 - self.beta:
return False
a_hat = self.expand_a(rho)
mu = dsa_H(dsa_H(pk, 64) + m_prime, 64)
c_hat = ntt_dsa(sample_in_ball(c_tilde, self.tau))
z_hat = [ntt_dsa(p) for p in z]
t1_hat = [ntt_dsa([(c << D_DROP) % Q_DSA for c in p]) for p in t1]
w1 = []
for i in range(self.k):
acc = [0] * 256
for j in range(self.l):
acc = [(x + y) % Q_DSA for x, y in
zip(acc, pointwise_dsa(a_hat[i][j], z_hat[j]))]
acc = [(x - y) % Q_DSA for x, y in
zip(acc, pointwise_dsa(c_hat, t1_hat[i]))]
wa = INTT_dsa(acc)
w1.append([use_hint(h[i][j], wa[j], self.gamma2) for j in range(256)])
return c_tilde == dsa_H(mu + self.w1_encode(w1), self.c_tilde_bytes)
def verify(self, pk, message, sig, ctx=b""):
if len(ctx) > 255:
return False
return self.verify_internal(pk, bytes([0, len(ctx)]) + ctx + message, sig)
print("ML-DSA implemented.")
ML-DSA implemented. PASSED (1.5 ms)
def test_mldsa(verbose=True):
checks = []
def ok(name, cond):
checks.append((name, bool(cond)))
if verbose:
print(" %-52s %s" % (name, "ok" if cond else "FAILED"))
ok("zeta has order 512 mod q",
pow(1753, 512, Q_DSA) == 1 and pow(1753, 256, Q_DSA) == Q_DSA - 1)
rng = random.Random(2)
a = [rng.randrange(Q_DSA) for _ in range(256)]
b = [rng.randrange(Q_DSA) for _ in range(256)]
ok("NTT round trip", INTT_dsa(ntt_dsa(a)) == a)
ok("NTT multiply == schoolbook",
INTT_dsa(pointwise_dsa(ntt_dsa(a), ntt_dsa(b))) == poly_mul_schoolbook(a, b, Q_DSA))
msg = b"embedded world north america"
for lvl, (pk_n, sk_n, sig_n) in ((44, (1312, 2560, 2420)),
(65, (1952, 4032, 3309)),
(87, (2592, 4896, 4627))):
d = MLDSA(lvl)
pk, sk = d.keygen()
sig = d.sign(sk, msg)
ok("ML-DSA-%d sizes match FIPS 204" % lvl,
(len(pk), len(sk), len(sig)) == (pk_n, sk_n, sig_n))
ok("ML-DSA-%d genuine signature verifies" % lvl, d.verify(pk, msg, sig))
ok("ML-DSA-%d altered message is rejected" % lvl,
not d.verify(pk, msg + b"!", sig))
bad = bytearray(sig); bad[len(sig) // 2] ^= 0x01
ok("ML-DSA-%d altered signature is rejected" % lvl,
not d.verify(pk, msg, bytes(bad)))
pk2, _ = d.keygen()
ok("ML-DSA-%d wrong public key is rejected" % lvl,
not d.verify(pk2, msg, sig))
d = MLDSA(65)
pk, sk = d.keygen()
s1 = d.sign(sk, msg, deterministic=True)
s2 = d.sign(sk, msg, deterministic=True)
s3 = d.sign(sk, msg)
ok("deterministic mode is reproducible", s1 == s2)
ok("hedged mode gives a fresh signature", s1 != s3)
ok("both modes verify", d.verify(pk, msg, s1) and d.verify(pk, msg, s3))
ok("context string is bound into the signature",
d.verify(pk, msg, d.sign(sk, msg, ctx=b"A"), ctx=b"A")
and not d.verify(pk, msg, d.sign(sk, msg, ctx=b"A"), ctx=b"B"))
ok("truncated signature is rejected", not d.verify(pk, msg, s1[:-1]))
n_ok = sum(1 for _, c in checks if c)
print("\n%d of %d checks passed." % (n_ok, len(checks)))
return n_ok == len(checks)
assert test_mldsa()
zeta has order 512 mod q ok NTT round trip ok NTT multiply == schoolbook ok ML-DSA-44 sizes match FIPS 204 ok ML-DSA-44 genuine signature verifies ok ML-DSA-44 altered message is rejected ok ML-DSA-44 altered signature is rejected ok ML-DSA-44 wrong public key is rejected ok ML-DSA-65 sizes match FIPS 204 ok ML-DSA-65 genuine signature verifies ok ML-DSA-65 altered message is rejected ok ML-DSA-65 altered signature is rejected ok ML-DSA-65 wrong public key is rejected ok ML-DSA-87 sizes match FIPS 204 ok ML-DSA-87 genuine signature verifies ok ML-DSA-87 altered message is rejected ok ML-DSA-87 altered signature is rejected ok ML-DSA-87 wrong public key is rejected ok deterministic mode is reproducible ok hedged mode gives a fresh signature ok both modes verify ok context string is bound into the signature ok truncated signature is rejected ok 23 of 23 checks passed. PASSED (324.6 ms)
# ---- the narrative demo: intermediate values for ML-DSA (s1, s2, y, w, c, z, h) -----
# We step through ML-DSA-65 with intermediate mathematical objects from the slides displayed.
dsa = MLDSA(65)
keccak_reset()
def signed_dsa(c, q=Q_DSA):
return c if c <= q // 2 else c - q
def fmt_dsa(p, n=8, as_signed=True):
vals = [signed_dsa(c) if as_signed else c for c in p[:n]]
return "[" + ", ".join(f"{x:+d}" if as_signed else str(x) for x in vals) + f", ... (total {len(p)} coeffs)]"
print("=" * 78)
print("1. ALICE KEY GENERATION (FIPS 204 Section 5.1 / ML-DSA.KeyGen)")
print("=" * 78)
xi = os.urandom(32)
seed = dsa_H(xi + bytes([dsa.k, dsa.l]), 128)
rho, rho_p, key = seed[:32], seed[32:96], seed[96:128]
a_hat = dsa.expand_a(rho)
s1, s2 = dsa.expand_s(rho_p)
pk, sk = dsa.keygen_internal(xi)
print(f"[*] Parameters: ML-DSA-65 (k={dsa.k}, l={dsa.l}, eta={dsa.eta}, q={Q_DSA})")
print(f"[*] Seed rho (32 bytes): {rho.hex()[:24]}...")
print(f" Matrix A-hat: shape {len(a_hat)} rows x {len(a_hat[0])} cols of NTT-domain polynomials (each 256 coeffs in Z_q)")
print(f"[*] Secret vector s1 in R_q^l:")
print(f" Shape: column vector of {len(s1)} polynomials x {len(s1[0])} small coeffs in [-eta, eta] = [-{dsa.eta}, +{dsa.eta}]")
for i in range(dsa.l):
print(f" s1[{i}][:8] = {fmt_dsa(s1[i])} range: [{min(signed_dsa(x) for x in s1[i])}, {max(signed_dsa(x) for x in s1[i])}]")
print(f"[*] Secret vector s2 in R_q^k:")
print(f" Shape: column vector of {len(s2)} polynomials x {len(s2[0])} small coeffs in [-eta, eta] = [-{dsa.eta}, +{dsa.eta}]")
for i in range(dsa.k):
print(f" s2[{i}][:8] = {fmt_dsa(s2[i])} range: [{min(signed_dsa(x) for x in s2[i])}, {max(signed_dsa(x) for x in s2[i])}]")
print(f"[*] Verification key pk: {len(pk)} bytes (t1 packed + rho)")
print(f"[*] Signing key sk: {len(sk)} bytes (rho + key + tr + s1 + s2 + t0)\n")
print("=" * 78)
print("2. SIGNING WITH REJECTION SAMPLING (FIPS 204 Section 5.2 / ML-DSA.Sign)")
print("=" * 78)
firmware = b"FIRMWARE v2.4.1 " + bytes(range(256)) * 4
print(f"Alice signs a {len(firmware)}-byte firmware image with context ctx=b'fw-update'.\n")
# Step through signing internals to capture accepted and rejected iteration values
rho_sk, key_sk, tr_sk, s1_dec, s2_dec, t0_dec = dsa.sk_decode(sk)
s1_hat = [ntt_dsa(p) for p in s1_dec]
s2_hat = [ntt_dsa(p) for p in s2_dec]
t0_hat = [ntt_dsa(p) for p in t0_dec]
m_prime = bytes([0, len(b"fw-update")]) + b"fw-update" + firmware
mu = dsa_H(tr_sk + m_prime, 64)
rho_pp = dsa_H(key_sk + bytes(32) + mu, 64)
kappa, attempts = 0, 0
diagnostics = []
while True:
attempts += 1
y = dsa.expand_mask(rho_pp, kappa)
kappa += dsa.l
y_hat = [ntt_dsa(p) for p in y]
w = []
for i in range(dsa.k):
acc = [0] * 256
for j in range(dsa.l):
acc = [(x + yy) % Q_DSA for x, yy in
zip(acc, pointwise_dsa(a_hat[i][j], y_hat[j]))]
w.append(INTT_dsa(acc))
w1 = [[high_bits(c, dsa.gamma2) for c in p] for p in w]
c_tilde = dsa_H(mu + dsa.w1_encode(w1), dsa.c_tilde_bytes)
c_poly = sample_in_ball(c_tilde, dsa.tau)
c_hat = ntt_dsa(c_poly)
cs1 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in s1_hat]
cs2 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in s2_hat]
z = [[(a + b) % Q_DSA for a, b in zip(p, q)] for p, q in zip(y, cs1)]
wcs2 = [[(a - b) % Q_DSA for a, b in zip(p, q)] for p, q in zip(w, cs2)]
r0 = [[low_bits(c, dsa.gamma2) for c in p] for p in wcs2]
norm_z = vec_norm(z)
norm_r0 = max(max(abs(c) for c in p) for p in r0)
# Check rejection conditions
if norm_z >= dsa.gamma1 - dsa.beta:
diagnostics.append((attempts, f"REJECTED: ||z||_inf = {norm_z} >= {dsa.gamma1 - dsa.beta} (gamma1 - beta)"))
continue
if norm_r0 >= dsa.gamma2 - dsa.beta:
diagnostics.append((attempts, f"REJECTED: ||r0||_inf = {norm_r0} >= {dsa.gamma2 - dsa.beta} (gamma2 - beta)"))
continue
ct0 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in t0_hat]
norm_ct0 = vec_norm(ct0)
if norm_ct0 >= dsa.gamma2:
diagnostics.append((attempts, f"REJECTED: ||c*t0||_inf = {norm_ct0} >= {dsa.gamma2}"))
continue
h = [[make_hint((-ct0[i][j]) % Q_DSA,
(wcs2[i][j] + ct0[i][j]) % Q_DSA, dsa.gamma2)
for j in range(256)] for i in range(dsa.k)]
hint_count = sum(sum(row) for row in h)
if hint_count > dsa.omega:
diagnostics.append((attempts, f"REJECTED: hint count = {hint_count} > {dsa.omega} (omega)"))
continue
diagnostics.append((attempts, f"ACCEPTED: ||z||_inf={norm_z} < {dsa.gamma1 - dsa.beta}, ||r0||_inf={norm_r0} < {dsa.gamma2 - dsa.beta}, hints={hint_count} <= {dsa.omega}"))
sig = dsa.sig_encode(c_tilde, z, h)
break
print("[*] [Item 3] Rejection Sampling Diagnostics Per Attempt:")
for att_num, diag in diagnostics:
status = "-> PASS" if "ACCEPTED" in diag else " FAIL"
print(f" Attempt {att_num}: {status} | {diag}")
print()
# ---- [ITEM 4] HIGHBITS / LOWBITS SPLIT (w = w1 * 2*gamma2 + w0) ----------
print("[*] [Item 4] HighBits / LowBits Decomposition (w = A*y mod q):")
print(f" Formula: w = w1 * (2*gamma2) + w0, where 2*gamma2 = {2*dsa.gamma2}")
print(" First 4 coefficients of w[0]:")
for j in range(4):
w_raw = w[0][j]
w1_val = high_bits(w_raw, dsa.gamma2)
w0_val = low_bits(w_raw, dsa.gamma2)
print(f" w[0][{j}] = {w_raw:7d} -> HighBits (w1) = {w1_val:2d}, LowBits (w0) = {w0_val:+6d} [check: {w1_val}*{2*dsa.gamma2} + ({w0_val}) = {w1_val*2*dsa.gamma2 + w0_val}]")
print()
print(f"[*] Ephemeral masking vector y in R_q^l:")
print(f" Shape: column vector of {len(y)} polynomials x {len(y[0])} coefficients in [-gamma1, gamma1]")
for i in range(dsa.l):
print(f" y[{i}][:8] = {fmt_dsa(y[i])}")
print(f"[*] Challenge hash c_tilde: {c_tilde.hex()[:24]}... ({len(c_tilde)} bytes)")
print(f" Challenge polynomial c(X) = SampleInBall(c_tilde): 256 coefficients with exactly tau={dsa.tau} non-zeros in {{-1, +1}}")
print(f" c[:16] = {fmt_dsa(c_poly[:16])}")
print(f"[*] Signature component z = y + c*s1 mod q (leaks no secret because bounded by gamma1 - beta):")
print(f" Shape: column vector of {len(z)} polynomials x {len(z[0])} coefficients")
for i in range(dsa.l):
print(f" z[{i}][:8] = {fmt_dsa(z[i])} ||z[{i}]||_inf = {max(abs(signed_dsa(x)) for x in z[i])} < {dsa.gamma1 - dsa.beta}")
print(f"[*] Hint bitmatrix h: shape {len(h)} rows x {len(h[0])} cols, total {hint_count} bits set (bound omega={dsa.omega})")
print(f"[*] Encoded signature: {len(sig)} bytes (c_tilde: {dsa.c_tilde_bytes} B + z: {32*dsa.z_bits*dsa.l} B + hints: {dsa.omega + dsa.k} B)\n")
print("=" * 78)
print("3. VERIFICATION (FIPS 204 Section 5.3 / ML-DSA.Verify)")
print("=" * 78)
print("The device verifies:")
print(" genuine firmware image: ", dsa.verify(pk, firmware, sig, ctx=b'fw-update'))
print(" one byte flipped: ", dsa.verify(pk, firmware[:-1] + b'\x00', sig, ctx=b'fw-update'))
print(" right sig, wrong context:", dsa.verify(pk, firmware, sig, ctx=b'boot'))
print()
print("Keccak for one keygen + one sign + three verifies: %d calls, ~%d permutations" % (KECCAK["calls"], KECCAK["perms"]))
# Rich Summary Cards for ML-DSA-65
show_card("ML-DSA-65 Key Generation Summary",
[("Seed rho (32 bytes)", f"{rho.hex()[:16]}..."),
("Secret Vector s1[0][:8]", chips([signed_dsa(x) for x in s1[0][:8]])),
("Secret Vector s2[0][:8]", chips([signed_dsa(x) for x in s2[0][:8]])),
("Verification Key pk Size", f"{len(pk)} bytes (t1 packed + rho)"),
("Signing Key sk Size", f"{len(sk)} bytes (rho + key + tr + s1 + s2 + t0)")],
category='secret', badge_text="KeyGen Stage")
show_card(f"ML-DSA-65 Signature Execution (Accepted on Attempt #{attempts})",
[("Signing Attempts Required", f"{attempts} attempt(s)"),
("Challenge Hash c_tilde", f"{c_tilde.hex()[:24]}..."),
("Signature Vector z[0][:8]", chips([signed_dsa(x) for x in z[0][:8]])),
("Norm Bound ||z||_∞", f"{norm_z:,} < {dsa.gamma1 - dsa.beta:,} " + badge("Norm Satisfied", "verified")),
("Hint Bits Set in Matrix h", f"{hint_count} / {dsa.omega} " + badge("Hints Valid", "verified")),
("Encoded Signature Size", f"{len(sig)} bytes (c_tilde: {dsa.c_tilde_bytes}B, z: 3200B, h: {dsa.omega + dsa.k}B)")],
category='verified', badge_text="Signature Generated")
# Complete 4-Step Verification from Slide 64 / FIPS 204 Section 5.3
# Step 1: Decode and unpack
c_tilde_dec, z_dec, h_dec = dsa.sig_decode(sig)
rho_dec, t1_dec = dsa.pk_decode(pk)
norm_z_verify = vec_norm(z_dec)
cond1_norm = (norm_z_verify < dsa.gamma1 - dsa.beta)
# Step 2: Weight of hint matrix h <= omega
hint_wt = sum(sum(row) for row in h_dec)
cond2_hint_wt = (hint_wt <= dsa.omega)
# Step 3: Reconstruction of w'1 via UseHint(h, Az - c*t1*2^d, 2*gamma2)
a_hat_v = dsa.expand_a(rho_dec)
mu_v = dsa_H(dsa_H(pk, 64) + bytes([0, len(b"fw-update")]) + b"fw-update" + firmware, 64)
c_hat_v = ntt_dsa(sample_in_ball(c_tilde_dec, dsa.tau))
z_hat_v = [ntt_dsa(p) for p in z_dec]
t1_hat_v = [ntt_dsa([(c << D_DROP) % Q_DSA for c in p]) for p in t1_dec]
w1_prime = []
for i in range(dsa.k):
acc = [0] * 256
for j in range(dsa.l):
acc = [(x + y) % Q_DSA for x, y in zip(acc, pointwise_dsa(a_hat_v[i][j], z_hat_v[j]))]
acc = [(x - y) % Q_DSA for x, y in zip(acc, pointwise_dsa(c_hat_v, t1_hat_v[i]))]
wa = INTT_dsa(acc)
w1_prime.append([use_hint(h_dec[i][j], wa[j], dsa.gamma2) for j in range(256)])
# Check condition 3: High-bits match exactly
cond3_w1_match = (w1_prime == w1)
# Step 4: Recomputed hash commitment c' == c_tilde
c_prime = dsa_H(mu_v + dsa.w1_encode(w1_prime), dsa.c_tilde_bytes)
cond4_hash_match = (c_prime == c_tilde_dec)
show_card("ML-DSA-65 Verification: 4 Architectural Checks (Slide 64 / FIPS 204)",
[("Condition 1: Response Vector Norm", f"||z||_∞ = {norm_z_verify:,} < {dsa.gamma1 - dsa.beta:,} " + badge("Condition 1 Pass", "verified")),
("Condition 2: Hint Matrix Weight", f"wt(h) = {hint_wt} <= {dsa.omega} (bound omega) " + badge("Condition 2 Pass", "verified")),
("Condition 3: High-Bits Recovery", f"w'1 == w1 in all {dsa.k} polynomials " + badge("Condition 3 Pass", "verified")),
("Condition 4: Challenge Hash Match", f"c' == c~ ({c_prime.hex()[:16]}... == {c_tilde_dec.hex()[:16]}...) " + badge("Condition 4 Pass", "verified")),
("Overall Verification Verdict", "All 4 Conditions Evaluated True -> ACCEPT (⊤) " + badge("Cryptographically Valid", "verified"))],
category='teal', badge_text="4-Condition Verification",
footer="Reference: Slide 64 (ML-DSA Sign Part 2) & FIPS 204 Algorithm 3. Verification requires all 4 conditions to evaluate true.")
v1 = dsa.verify(pk, firmware, sig, ctx=b'fw-update')
v2 = dsa.verify(pk, firmware[:-1] + b'\x00', sig, ctx=b'fw-update')
v3 = dsa.verify(pk, firmware, sig, ctx=b'boot')
show_card("ML-DSA-65 Tamper Resistance Tests",
[("Genuine Firmware Image", f"{v1} " + badge("Accept Signature", "verified")),
("One Byte Flipped", f"{v2} " + badge("Reject Signature (Tampered)", "reject")),
("Wrong Execution Context (ctx=b'boot')", f"{v3} " + badge("Reject Signature (Wrong Context)", "reject"))],
category='teal', badge_text="Tamper Tests")
==============================================================================
1. ALICE KEY GENERATION (FIPS 204 Section 5.1 / ML-DSA.KeyGen)
==============================================================================
[*] Parameters: ML-DSA-65 (k=6, l=5, eta=4, q=8380417)
[*] Seed rho (32 bytes): 420ed13318e20bc83f9ffcd2...
Matrix A-hat: shape 6 rows x 5 cols of NTT-domain polynomials (each 256 coeffs in Z_q)
[*] Secret vector s1 in R_q^l:
Shape: column vector of 5 polynomials x 256 small coeffs in [-eta, eta] = [-4, +4]
s1[0][:8] = [+3, +2, -1, -4, +4, +2, -3, -3, ... (total 256 coeffs)] range: [-4, 4]
s1[1][:8] = [-4, +0, -1, +1, +0, -4, +0, +2, ... (total 256 coeffs)] range: [-4, 4]
s1[2][:8] = [+1, -3, +2, +3, +4, +3, -1, -3, ... (total 256 coeffs)] range: [-4, 4]
s1[3][:8] = [+4, +4, +4, -2, +4, -1, +4, -1, ... (total 256 coeffs)] range: [-4, 4]
s1[4][:8] = [-1, +1, -4, +1, +0, +4, -1, +0, ... (total 256 coeffs)] range: [-4, 4]
[*] Secret vector s2 in R_q^k:
Shape: column vector of 6 polynomials x 256 small coeffs in [-eta, eta] = [-4, +4]
s2[0][:8] = [-2, +3, +2, -1, +3, -2, +1, -4, ... (total 256 coeffs)] range: [-4, 4]
s2[1][:8] = [+3, -3, -2, +4, -2, +1, +1, -2, ... (total 256 coeffs)] range: [-4, 4]
s2[2][:8] = [+2, -1, +1, +2, +4, +0, +4, +1, ... (total 256 coeffs)] range: [-4, 4]
s2[3][:8] = [+0, -3, +2, -1, -4, -4, -2, -4, ... (total 256 coeffs)] range: [-4, 4]
s2[4][:8] = [+0, +4, -4, -1, -4, +1, -3, +4, ... (total 256 coeffs)] range: [-4, 4]
s2[5][:8] = [+0, +1, -3, -2, +0, -2, -2, -1, ... (total 256 coeffs)] range: [-4, 4]
[*] Verification key pk: 1952 bytes (t1 packed + rho)
[*] Signing key sk: 4032 bytes (rho + key + tr + s1 + s2 + t0)
==============================================================================
2. SIGNING WITH REJECTION SAMPLING (FIPS 204 Section 5.2 / ML-DSA.Sign)
==============================================================================
Alice signs a 1040-byte firmware image with context ctx=b'fw-update'.
[*] [Item 3] Rejection Sampling Diagnostics Per Attempt:
Attempt 1: FAIL | REJECTED: ||z||_inf = 524161 >= 524092 (gamma1 - beta)
Attempt 2: FAIL | REJECTED: ||r0||_inf = 261820 >= 261692 (gamma2 - beta)
Attempt 3: FAIL | REJECTED: ||z||_inf = 524254 >= 524092 (gamma1 - beta)
Attempt 4: FAIL | REJECTED: ||z||_inf = 524250 >= 524092 (gamma1 - beta)
Attempt 5: -> PASS | ACCEPTED: ||z||_inf=523977 < 524092, ||r0||_inf=261277 < 261692, hints=36 <= 55
[*] [Item 4] HighBits / LowBits Decomposition (w = A*y mod q):
Formula: w = w1 * (2*gamma2) + w0, where 2*gamma2 = 523776
First 4 coefficients of w[0]:
w[0][0] = 1448721 -> HighBits (w1) = 3, LowBits (w0) = -122607 [check: 3*523776 + (-122607) = 1448721]
w[0][1] = 5825246 -> HighBits (w1) = 11, LowBits (w0) = +63710 [check: 11*523776 + (63710) = 5825246]
w[0][2] = 6984047 -> HighBits (w1) = 13, LowBits (w0) = +174959 [check: 13*523776 + (174959) = 6984047]
w[0][3] = 4922692 -> HighBits (w1) = 9, LowBits (w0) = +208708 [check: 9*523776 + (208708) = 4922692]
[*] Ephemeral masking vector y in R_q^l:
Shape: column vector of 5 polynomials x 256 coefficients in [-gamma1, gamma1]
y[0][:8] = [+66821, +519719, +484170, +105232, +208961, +284604, -353852, -52751, ... (total 256 coeffs)]
y[1][:8] = [-160810, +415041, +71414, +512783, +504733, -125564, +313103, -75894, ... (total 256 coeffs)]
y[2][:8] = [+191056, -187420, -424086, -93309, -54276, -28356, -24054, -281526, ... (total 256 coeffs)]
y[3][:8] = [+149499, -287151, -116631, -161921, +520844, -426939, +475070, -456199, ... (total 256 coeffs)]
y[4][:8] = [+457772, +185606, +332209, +173403, +428479, -14778, -89548, -93000, ... (total 256 coeffs)]
[*] Challenge hash c_tilde: d9428142022da25d1076ac60... (48 bytes)
Challenge polynomial c(X) = SampleInBall(c_tilde): 256 coefficients with exactly tau=49 non-zeros in {-1, +1}
c[:16] = [+0, +1, +0, +0, +0, +0, +0, +0, ... (total 16 coeffs)]
[*] Signature component z = y + c*s1 mod q (leaks no secret because bounded by gamma1 - beta):
Shape: column vector of 5 polynomials x 256 coefficients
z[0][:8] = [+66823, +519715, +484182, +105260, +208959, +284616, -353888, -52787, ... (total 256 coeffs)] ||z[0]||_inf = 523214 < 524092
z[1][:8] = [-160785, +415026, +71402, +512764, +504736, -125578, +313107, -75880, ... (total 256 coeffs)] ||z[1]||_inf = 523079 < 524092
z[2][:8] = [+191047, -187415, -424096, -93319, -54286, -28333, -24040, -281483, ... (total 256 coeffs)] ||z[2]||_inf = 523977 < 524092
z[3][:8] = [+149485, -287112, -116633, -161909, +520862, -426917, +475071, -456185, ... (total 256 coeffs)] ||z[3]||_inf = 522173 < 524092
z[4][:8] = [+457766, +185596, +332196, +173404, +428505, -14811, -89542, -93009, ... (total 256 coeffs)] ||z[4]||_inf = 523564 < 524092
[*] Hint bitmatrix h: shape 6 rows x 256 cols, total 36 bits set (bound omega=55)
[*] Encoded signature: 3309 bytes (c_tilde: 48 B + z: 3200 B + hints: 61 B)
==============================================================================
3. VERIFICATION (FIPS 204 Section 5.3 / ML-DSA.Verify)
==============================================================================
The device verifies:
genuine firmware image: True
one byte flipped: False
right sig, wrong context: False
Keccak for one keygen + one sign + three verifies: 374 calls, ~1300 permutations
PASSED (64.3 ms)
# ---- EACH TIME WE GENERATE A NEW SIGNATURE (REJECTION SAMPLING) -----------
# A single message signed with full visibility into each retry loop attempt until acceptance.
def sign_with_trace(dsa, sk, message, ctx=b""):
m_prime = bytes([0, len(ctx)]) + ctx + message
rho, key, tr, s1, s2, t0 = dsa.sk_decode(sk)
s1_hat = [ntt_dsa(p) for p in s1]
s2_hat = [ntt_dsa(p) for p in s2]
t0_hat = [ntt_dsa(p) for p in t0]
a_hat = dsa.expand_a(rho)
mu = dsa_H(tr + m_prime, 64)
# Try random nonces until we find one that exhibits rejections (for rich visual demonstration)
for trial_seed in range(100):
rnd = sha3_256(b"ewna-seed" + bytes([trial_seed])).digest()
rho_pp = dsa_H(key + rnd + mu, 64)
attempts = []
kappa = 0
while True:
attempt_idx = len(attempts) + 1
y = dsa.expand_mask(rho_pp, kappa)
kappa += dsa.l
y_hat = [ntt_dsa(p) for p in y]
w = []
for i in range(dsa.k):
acc = [0] * 256
for j in range(dsa.l):
acc = [(x + yy) % Q_DSA for x, yy in zip(acc, pointwise_dsa(a_hat[i][j], y_hat[j]))]
w.append(INTT_dsa(acc))
w1 = [[high_bits(c, dsa.gamma2) for c in p] for p in w]
c_tilde = dsa_H(mu + dsa.w1_encode(w1), dsa.c_tilde_bytes)
c_hat = ntt_dsa(sample_in_ball(c_tilde, dsa.tau))
cs1 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in s1_hat]
cs2 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in s2_hat]
z = [[(a + b) % Q_DSA for a, b in zip(p, q)] for p, q in zip(y, cs1)]
wcs2 = [[(a - b) % Q_DSA for a, b in zip(p, q)] for p, q in zip(w, cs2)]
r0 = [[low_bits(c, dsa.gamma2) for c in p] for p in wcs2]
norm_z = vec_norm(z)
max_r0 = max(max(abs(c) for c in p) for p in r0)
if norm_z >= dsa.gamma1 - dsa.beta:
attempts.append({"attempt": attempt_idx, "status": "REJECTED", "cause": "z_norm_overflow",
"details": f"||z||_∞ = {norm_z:,} >= bound {dsa.gamma1 - dsa.beta:,} (would leak secret s1 bits)",
"c_tilde": c_tilde, "z": z, "r0_max": max_r0})
continue
if max_r0 >= dsa.gamma2 - dsa.beta:
attempts.append({"attempt": attempt_idx, "status": "REJECTED", "cause": "r0_norm_overflow",
"details": f"||r0||_∞ = {max_r0:,} >= bound {dsa.gamma2 - dsa.beta:,} (low bits would fail high-bit hint recovery)",
"c_tilde": c_tilde, "z": z, "r0_max": max_r0})
continue
ct0 = [INTT_dsa(pointwise_dsa(c_hat, p)) for p in t0_hat]
if vec_norm(ct0) >= dsa.gamma2:
attempts.append({"attempt": attempt_idx, "status": "REJECTED", "cause": "ct0_overflow",
"details": f"||c*t0||_∞ = {vec_norm(ct0):,} >= bound {dsa.gamma2:,}",
"c_tilde": c_tilde, "z": z, "r0_max": max_r0})
continue
h = [[make_hint((-ct0[i][j]) % Q_DSA, (wcs2[i][j] + ct0[i][j]) % Q_DSA, dsa.gamma2)
for j in range(256)] for i in range(dsa.k)]
hint_count = sum(sum(row) for row in h)
if hint_count > dsa.omega:
attempts.append({"attempt": attempt_idx, "status": "REJECTED", "cause": "hint_overflow",
"details": f"Hint count = {hint_count} > max omega {dsa.omega}",
"c_tilde": c_tilde, "z": z, "r0_max": max_r0})
continue
sig = dsa.sig_encode(c_tilde, z, h)
attempts.append({"attempt": attempt_idx, "status": "ACCEPTED", "cause": None,
"details": f"All norm bounds and hint weights satisfied!",
"c_tilde": c_tilde, "z": z, "norm_z": norm_z, "r0_max": max_r0,
"hints": hint_count, "sig": sig})
break
if len(attempts) >= 3:
return attempts
dsa = MLDSA(65)
pk, sk = dsa.keygen()
target_firmware = b"FIRMWARE v3.2.0-PATCH-2026"
attempts_log = sign_with_trace(dsa, sk, target_firmware, ctx=b"secure-boot")
for step in attempts_log:
att_num = step["attempt"]
if step["status"] == "REJECTED":
show_card(f"ML-DSA-65 Signing Attempt #{att_num} [REJECTED]",
[("Input Message", f"'{target_firmware.decode()}'"),
("Attempt Verdict", badge("Rejected (Retry Next Candidate)", "reject")),
("Rejection Reason", step["details"]),
("Candidate Mask y / z[0][:6]", chips([signed_dsa(x) for x in step["z"][0][:6]])),
("Challenge Hash c_tilde", f"{step['c_tilde'].hex()[:24]}..."),
("Security Rationale", "If this signature were published, an adversary collecting samples could compute the private key s1/s2.")],
category='reject', badge_text=f"Attempt #{att_num} Aborted",
footer="Fiat-Shamir with Aborts: Discard candidate and re-sample mask vector y without incrementing any public state.")
else:
valid = dsa.verify(pk, target_firmware, step["sig"], ctx=b"secure-boot")
show_card(f"ML-DSA-65 Signing Attempt #{att_num} [ACCEPTED & VERIFIED]",
[("Input Message", f"'{target_firmware.decode()}'"),
("Attempt Verdict", badge(f"Accepted on Attempt #{att_num}!", "verified")),
("Norm Check ||z||_∞", f"{step['norm_z']:,} < {dsa.gamma1 - dsa.beta:,} " + badge("Strictly Bounded", "verified")),
("Low-Bits Check ||r0||_∞", f"{step['r0_max']:,} < {dsa.gamma2 - dsa.beta:,} " + badge("Hints Reliable", "verified")),
("Valid Hints Set in h", f"{step['hints']} / {dsa.omega} " + badge("Hints Fit Packing", "verified")),
("Final Signature Vector z[0][:6]", chips([signed_dsa(x) for x in step["z"][0][:6]])),
("Verification Check", f"verify(pk, msg, sig) == {valid} " + badge("Cryptographically Valid", "verified"))],
category='verified', badge_text=f"Valid Signature (Attempt #{att_num})",
footer=f"Total loop iterations required: {att_num}. Secret key privacy mathematically preserved.")
PASSED (53.2 ms)
# ---- REJECTION STATISTICS ACROSS PARAMETER SETS (NUMERICAL TABLE) --------
def rejection_stats(level=65, n=40, seed=None):
d = MLDSA(level)
pk, sk = d.keygen()
tries, reasons = [], {}
for i in range(n):
m = b"message %d" % i
sig, t, why = d.sign(sk, m, stats=True)
assert d.verify(pk, m, sig)
tries.append(t)
for r in why:
reasons[r] = reasons.get(r, 0) + 1
return tries, reasons
EXPECTED = {44: 4.25, 65: 5.10, 87: 3.85}
stats_rows = []
for lvl in (44, 65, 87):
tries, reasons = rejection_stats(lvl, 40)
mean_t = statistics.mean(tries)
med_t = statistics.median(tries)
min_t = min(tries)
max_t = max(tries)
total_rej = sum(reasons.values())
z_pct = reasons.get("z too large", 0) / max(1, total_rej) * 100
r0_pct = reasons.get("r0 too large", 0) / max(1, total_rej) * 100
h_pct = (reasons.get("too many hints", 0) + reasons.get("c*t0 too large", 0)) / max(1, total_rej) * 100
stats_rows.append([
f"ML-DSA-{lvl}",
f"{mean_t:.2f} (spec: {EXPECTED[lvl]:.2f})",
f"{med_t:.1f}",
f"{min_t} / {max_t}",
f"{z_pct:4.1f} %",
f"{r0_pct:4.1f} %",
f"{h_pct:4.1f} %"
])
show_table(["Security Level", "Mean Attempts", "Median", "Min / Max", "z Rejection %", "r0 Rejection %", "Hint Rejection %"],
stats_rows, title="ML-DSA Rejection Loop Dynamics across Parameter Sets", category='mauve')
PASSED (3280.4 ms)
| Security Level | Mean Attempts | Median | Min / Max | z Rejection % | r0 Rejection % | Hint Rejection % |
|---|---|---|---|---|---|---|
| ML-DSA-44 | 4.10 (spec: 4.25) | 3.0 | 1 / 16 | 58.1 % | 41.1 % | 0.8 % |
| ML-DSA-65 | 4.72 (spec: 5.10) | 4.0 | 1 / 15 | 55.0 % | 45.0 % | 0.0 % |
| ML-DSA-87 | 3.85 (spec: 3.85) | 3.0 | 1 / 14 | 48.2 % | 50.9 % | 0.9 % |
---
Now that Keccak and NTT are instrumented, convert the operation counts into estimated hardware cycles.
NTT_COUNT = {"kem": 0, "dsa": 0}
_ntt_kem, _INTT_kem = ntt_kem, INTT_kem
_ntt_dsa, _INTT_dsa = ntt_dsa, INTT_dsa
def ntt_kem(f):
NTT_COUNT["kem"] += 1; return _ntt_kem(f)
def INTT_kem(f):
NTT_COUNT["kem"] += 1; return _INTT_kem(f)
def ntt_dsa(f):
NTT_COUNT["dsa"] += 1; return _ntt_dsa(f)
def INTT_dsa(f):
NTT_COUNT["dsa"] += 1; return _INTT_dsa(f)
def profile(fn, kind):
keccak_reset()
NTT_COUNT["kem"] = NTT_COUNT["dsa"] = 0
t0 = time.perf_counter()
fn()
wall = time.perf_counter() - t0
return dict(perms=KECCAK["perms"], calls=KECCAK["calls"],
ntts=NTT_COUNT[kind], wall=wall,
detail=dict(KECCAK["detail"]))
kem = MLKEM(768)
ek, dk = kem.keygen()
_, ct = kem.encaps(ek)
dsa = MLDSA(65)
pk, sk = dsa.keygen()
msg = b"profile me"
sig = dsa.sign(sk, msg)
jobs = [
("ML-KEM-768 KeyGen", lambda: kem.keygen(), "kem"),
("ML-KEM-768 Encaps", lambda: kem.encaps(ek), "kem"),
("ML-KEM-768 Decaps", lambda: kem.decaps(dk, ct), "kem"),
("ML-DSA-65 KeyGen", lambda: dsa.keygen(), "dsa"),
("ML-DSA-65 Sign", lambda: dsa.sign(sk, msg), "dsa"),
("ML-DSA-65 Verify", lambda: dsa.verify(pk, msg, sig), "dsa"),
]
results = {}
rows = []
for name, fn, kind in jobs:
r = profile(fn, kind)
results[name] = r
rows.append([name, r["calls"], r["perms"], r["ntts"], "%.0f" % (1000 * r["wall"])])
table(["operation", "Keccak calls", "Keccak perms", "NTT/INTT calls", "ms here"], rows)
print()
print("A Keccak permutation is 24 rounds. A 256-point NTT is 1024 butterflies")
print("(896 for ML-KEM's 7 stages). Multiply those out for a hardware estimate.")
operation Keccak calls Keccak perms NTT/INTT calls ms here ----------------- ------------ ------------ -------------- ------- ML-KEM-768 KeyGen 17 43 6 2 ML-KEM-768 Encaps 18 44 7 3 ML-KEM-768 Decaps 19 53 11 3 ML-DSA-65 KeyGen 73 240 11 5 ML-DSA-65 Sign 90 326 115 21 ML-DSA-65 Verify 64 207 18 5 A Keccak permutation is 24 rounds. A 256-point NTT is 1024 butterflies (896 for ML-KEM's 7 stages). Multiply those out for a hardware estimate. PASSED (57.6 ms)
# Hardware Cycle Split Calculation
KECCAK_CYCLES_PER_PERM = 24 # 24 rounds per permutation
BUTTERFLY_CYCLES = 1 # 1 butterfly per cycle in dedicated DSP
hw_rows = []
for name, r in results.items():
stages = 7 if "KEM" in name else 8
ntt_cycles = r["ntts"] * (256 // 2) * stages * BUTTERFLY_CYCLES
kec_cycles = r["perms"] * KECCAK_CYCLES_PER_PERM
total = ntt_cycles + kec_cycles
hw_rows.append([name, f"{kec_cycles:,}", f"{ntt_cycles:,}", f"{total:,}",
f"{100 * kec_cycles / total:5.1f} %",
f"{100 * ntt_cycles / total:5.1f} %"])
show_table(["Operation", "Keccak (cyc)", "NTT (cyc)", "Total (cyc)", "Keccak %", "NTT %"],
hw_rows, title="Hardware Cycle Split Estimate (1 Butterfly Engine + 1-Round/Cycle Keccak)",
category='mauve', highlight_col=3)
print()
print("Compare that with the software split (Keccak 50-70 %) and you have the")
print("most useful result in this notebook: moving to hardware speeds Keccak up")
print("by roughly 500x and the NTT by roughly 7x, so the bottleneck *moves*.")
print("In software, accelerate Keccak. In hardware, buy butterflies and banks.")
Compare that with the software split (Keccak 50-70 %) and you have the most useful result in this notebook: moving to hardware speeds Keccak up by roughly 500x and the NTT by roughly 7x, so the bottleneck *moves*. In software, accelerate Keccak. In hardware, buy butterflies and banks. PASSED (0.2 ms)
| Operation | Keccak (cyc) | NTT (cyc) | Total (cyc) | Keccak % | NTT % |
|---|---|---|---|---|---|
| ML-KEM-768 KeyGen | 1,032 | 5,376 | 6,408 | 16.1 % | 83.9 % |
| ML-KEM-768 Encaps | 1,056 | 6,272 | 7,328 | 14.4 % | 85.6 % |
| ML-KEM-768 Decaps | 1,272 | 9,856 | 11,128 | 11.4 % | 88.6 % |
| ML-DSA-65 KeyGen | 5,760 | 11,264 | 17,024 | 33.8 % | 66.2 % |
| ML-DSA-65 Sign | 7,824 | 117,760 | 125,584 | 6.2 % | 93.8 % |
| ML-DSA-65 Verify | 4,968 | 18,432 | 23,400 | 21.2 % | 78.8 % |
---
Estimate latency and throughput based on butterfly parallelism and Keccak rounds per cycle.
def estimate(butterflies=1, keccak_rounds_per_cycle=1, clock_mhz=100,
kem_level=768, dsa_level=65, verbose=False):
kem_stages, dsa_stages = 7, 8
eff = min(butterflies, 16) ** 0.85 if butterflies > 1 else 1.0
r = results["ML-KEM-768 Encaps"]
ntt_c = int((r["ntts"] * 128 * 7) / (butterflies * eff))
kec_c = int(r["perms"] * (24 / keccak_rounds_per_cycle))
tot_c = ntt_c + kec_c
us = (tot_c / (clock_mhz * 1e6)) * 1e6
return kec_c, ntt_c, tot_c, us
scale_rows = []
for bf in [1, 2, 4, 8, 16, 32]:
_, _, tot_only_ntt, _ = estimate(butterflies=bf, keccak_rounds_per_cycle=1)
k_scal = min(bf, 24)
_, _, tot_both, lat = estimate(butterflies=bf, keccak_rounds_per_cycle=k_scal)
scale_rows.append([f"{bf} units", f"{tot_only_ntt:,} cyc", f"{tot_both:,} cyc", f"{lat:.1f} us"])
show_table(["Butterfly Engines", "Only NTT Scaled", "Both Scaled (Keccak + NTT)", "Latency @ 100 MHz"],
scale_rows, title="Parallel Butterfly Scaling for ML-KEM-768 Encapsulation", category='mauve', highlight_col=2)
PASSED (0.2 ms)
| Butterfly Engines | Only NTT Scaled | Both Scaled (Keccak + NTT) | Latency @ 100 MHz |
|---|---|---|---|
| 1 units | 7,328 cyc | 7,328 cyc | 73.3 us |
| 2 units | 2,795 cyc | 2,267 cyc | 22.7 us |
| 4 units | 1,538 cyc | 746 cyc | 7.5 us |
| 8 units | 1,189 cyc | 265 cyc | 2.6 us |
| 16 units | 1,093 cyc | 103 cyc | 1.0 us |
| 32 units | 1,074 cyc | 62 cyc | 0.6 us |
# Hardware Design Point Presets & Interactive Exploration
HAVE_WIDGETS = False
try:
import importlib
_widgets = importlib.import_module("ipywidgets")
interact = getattr(_widgets, "interact")
Dropdown = getattr(_widgets, "Dropdown")
IntSlider = getattr(_widgets, "IntSlider")
HAVE_WIDGETS = True
except (ImportError, ModuleNotFoundError, AttributeError, Exception):
HAVE_WIDGETS = False
def estimate_full(butterflies=1, keccak_rounds_per_cycle=1, clock_mhz=100,
kem_level=768, dsa_level=65, verbose=True):
kem_stages, dsa_stages = 7, 8
eff = min(butterflies, 16) ** 0.85 if butterflies > 1 else 1.0
rows = []
for name, r in results.items():
stages = kem_stages if "KEM" in name else dsa_stages
ntt_c = r["ntts"] * (256 // 2) * stages / eff
kec_c = r["perms"] * 24 / keccak_rounds_per_cycle
total = ntt_c + kec_c
rows.append([name, f"{int(kec_c):,}", f"{int(ntt_c):,}", f"{int(total):,}",
f"{total / (clock_mhz * 1000.0):.2f} ms",
f"{round(100 * kec_c / total)} %"])
if verbose:
header_text = f"Hardware Config: {butterflies} butterflies, {keccak_rounds_per_cycle} Keccak rd/cyc @ {clock_mhz} MHz (Eff speedup: {eff:.1f}x)"
show_table(["Operation", "Keccak cyc", "NTT cyc", "Total cyc", "Latency", "Keccak %"],
rows, title=header_text, category='mauve')
return rows
if HAVE_WIDGETS:
interact(estimate_full,
butterflies=Dropdown(options=[1, 2, 4, 8, 16, 32], value=1, description="Butterflies"),
keccak_rounds_per_cycle=Dropdown(options=[0.5, 1, 2, 4, 24], value=1, description="Keccak rd/cyc"),
clock_mhz=IntSlider(min=25, max=800, step=25, value=100, description="Clock MHz"),
kem_level=Dropdown(options=[512, 768, 1024], value=768),
dsa_level=Dropdown(options=[44, 65, 87], value=65),
verbose=True)
else:
preset_rows = []
for bf, kr, clk, label in ((1, 0.5, 50, "Tiny (Resource-Constrained IoT)"),
(2, 1, 100, "Balanced (Embedded SoC)"),
(8, 2, 300, "Fast (High-Throughput HSM)")):
rows_p = estimate_full(bf, kr, clk, verbose=False)
kem_enc = next(r for r in rows_p if "ML-KEM-768 Encaps" in r[0])
preset_rows.append([label, f"{bf} BF / {kr} rd/cyc", f"{clk} MHz", kem_enc[3], kem_enc[4]])
show_table(["Design Point Profile", "Hardware Resources", "Clock", "Total Cycles", "Latency"],
preset_rows, title="Representative Hardware Architecture Profiles (ML-KEM-768 Encaps)", category='teal')
PASSED (1.8 ms)
| Design Point Profile | Hardware Resources | Clock | Total Cycles | Latency |
|---|---|---|---|---|
| Tiny (Resource-Constrained IoT) | 1 BF / 0.5 rd/cyc | 50 MHz | 8,384 | 0.17 ms |
| Balanced (Embedded SoC) | 2 BF / 1 rd/cyc | 100 MHz | 4,535 | 0.05 ms |
| Fast (High-Throughput HSM) | 8 BF / 2 rd/cyc | 300 MHz | 1,598 | 0.01 ms |
---
The Fujisaki-Okamoto transform makes ML-KEM IND-CCA2 secure, provided the ciphertext comparison inside decapsulation is constant-time.
def leaky_compare(a, b):
if len(a) != len(b): return False
for x, y in zip(a, b):
if x != y: return False
return True
def safe_compare(a, b):
if len(a) != len(b): return False
diff = 0
for x, y in zip(a, b):
diff |= x ^ y
return diff == 0
buf_len = 2048
target_secret = b"\xaa" * buf_len
timing_rows = []
for match_len in [0, 256, 512, 1024, 1536, 2048]:
candidate = target_secret[:match_len] + b"\x00" * (buf_len - match_len)
t0 = time.perf_counter_ns()
for _ in range(500):
leaky_compare(target_secret, candidate)
t_leaky = (time.perf_counter_ns() - t0) / 500
t0 = time.perf_counter_ns()
for _ in range(500):
safe_compare(target_secret, candidate)
t_safe = (time.perf_counter_ns() - t0) / 500
timing_rows.append([f"{match_len} bytes", f"{t_leaky / 1000:6.2f} us", f"{t_safe / 1000:6.2f} us"])
show_table(["Matching Prefix Length", "Leaky Compare Time", "Safe Constant-Time Compare Time"],
timing_rows, title="Ciphertext Comparison Timing: Early-Exit Leak vs Constant-Time", category='reject', highlight_col=1)
PASSED (147.6 ms)
| Matching Prefix Length | Leaky Compare Time | Safe Constant-Time Compare Time |
|---|---|---|
| 0 bytes | 0.20 us | 41.66 us |
| 256 bytes | 2.98 us | 38.14 us |
| 512 bytes | 6.23 us | 38.67 us |
| 1024 bytes | 11.56 us | 38.52 us |
| 1536 bytes | 17.16 us | 38.38 us |
| 2048 bytes | 22.85 us | 38.32 us |
# What the leak is worth: live secret byte recovery with a timing oracle
QUERIES = {"n": 0}
secret_short = os.urandom(16)
def timing_oracle(candidate):
"""Returns how many bytes matched, which is what an early-exit stopwatch reveals."""
QUERIES["n"] += 1
for i, (x, y) in enumerate(zip(secret_short, candidate)):
if x != y:
return i
return len(secret_short)
QUERIES["n"] = 0
recovered = bytearray()
for pos in range(len(secret_short)):
for guess in range(256):
trial = bytes(recovered) + bytes([guess]) + bytes(len(secret_short) - pos - 1)
if timing_oracle(trial) > pos:
recovered.append(guess)
break
show_card("Timing Side-Channel Exploitation: Key Recovery",
[("Target Secret (16 bytes)", f"0x{secret_short.hex()}"),
("Recovered Secret via Timing", f"0x{bytes(recovered).hex()} " + badge("Exact Recovery", "reject")),
("Total Oracle Timing Queries", f"{QUERIES['n']} queries"),
("Brute Force Search Space", "2^128 ≈ 3.4 x 10^38 attempts"),
("Attack Complexity Reduction", f"Reduced from 2^128 to ~{QUERIES['n']} queries " + badge(f"Trivially Recovered in {QUERIES['n']} Steps", "reject")),
("Core Takeaway", "Early-exit ciphertext comparisons completely bypass IND-CCA2 security. Decaps must always execute in constant time.")],
category='reject', badge_text="Live Timing Exploit")
PASSED (1.5 ms)
dsa_timing = MLDSA(44)
pk_t, sk_t = dsa_timing.keygen()
signing_times = []
signing_attempts = []
for msg_idx in range(20):
t0 = time.perf_counter()
_, att_i, _ = dsa_timing.sign(sk_t, f"msg-{msg_idx}".encode(), stats=True)
dt = (time.perf_counter() - t0) * 1000
signing_times.append(dt)
signing_attempts.append(att_i)
show_card("ML-DSA Signing Timing Spread (Variable Loop Latency)",
[("Sample Message Count", f"{len(signing_times)} messages"),
("Attempt Distribution", f"Min: {min(signing_attempts)} | Median: {statistics.median(signing_attempts):.1f} | Max: {max(signing_attempts)}"),
("Execution Time Spread", f"Min: {min(signing_times):.1f} ms | Median: {statistics.median(signing_times):.1f} ms | Max: {max(signing_times):.1f} ms"),
("Timing Channel Implication", "Signing latency inherently varies due to rejection sampling; constant-time loops require dummy iterations.")],
category='mauve', badge_text="Signing Latency Spread")
PASSED (373.1 ms)
---
Key experiments to explore:
1. Break the sign flip: In poly_mul_schoolbook, change -= to +=. Run test_mlkem() and observe which check catches it first.
2. Increase noise bound: In toy_encrypt, set eta = 4 or eta = 6. Watch the bit error rate climb and see the exact coefficient where |w| >= q/4.
3. Implicit rejection bypass: In MLKEM.decaps_internal, remove the fallback and return key2 even on re-encryption mismatch. Notice how chosen-ciphertext attacks can now extract the private key.
4. Bypass ML-DSA rejection: In MLDSA.sign_internal, comment out the if vec_norm(z) >= self.gamma1 - self.beta: continue check. Collect 500 signatures and compute the average of z; observe how the mean reveals the secret key s1.