Hands-On Post-Quantum Cryptography: ML-KEM & ML-DSA

Complete Workshop Notebook with All Executed Outputs · FIPS 203 & FIPS 204 Reference Implementations

Post-Quantum Cryptography for Embedded Systems

Companion notebook, Embedded World North America 2026

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.

[1]:
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")
Out[1]:
PASSED (32.4 ms)
PQC Embedded Demonstration EnvironmentInitialized
Runtime Engine
Python 3.13.15
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

---

1. Lattices, Rings, and Toy KEM <span style="color:#DF8E1D">(Demo 1)</span>

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.

[2]:
# 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).")
Out[2]:
PASSED (0.2 ms)
Closest Vector Problem: Babai Rounding ComparisonCVP Decoding
Continuous Target Point
(2.7, 3.4)
Good Basis (Orthogonal)
b1=(1, 0), b2=(0, 1)
-> Decoded Integer Point
(3, 3) (combo: (3, 3))
-> Distance to Target
0.500 Exact Closest Point
Bad Basis (Skewed)
b1=(10, 11), b2=(11, 12)
-> Decoded Integer Point
(6, 7) (combo: (5, -4))
-> Distance to Target
4.884 Error 9.8x Larger
Takeaway: The public key is a 'bad basis' (hard CVP). The private key is a 'good basis' (easy noise removal).

1.1 Numbers on a clock face: Centred Modulo

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

[3]:
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")
Out[3]:
PASSED (0.2 ms)
Standard vs Centred Modulo Reduction (q = 13)
Value xStandard mod 13Centred mod 13 in (-6, +6]
1010-3
16640+0
16651+1
33280+0
00+0
77-6
Hardware Modular Arithmetic PropertiesZero DSP Multipliers
ML-KEM Modulus
q = 3329 = 2^8 * 13 + 1 (12 bits)
ML-DSA Modulus
q = 8380417 = 2^23 - 2^13 + 1 (23 bits)
ML-DSA Zero-DSP Trick
s * q = (s << 23) - (s << 13) + s
Verification for s=12345
s*q = 103456247865 == 103456247865 Verified: 2 shifts, 1 sub, 1 add

1.2 Polynomials in the ring: Negacyclic Multiplication

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

[4]:
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")
Out[4]:
PASSED (0.2 ms)
Negacyclic Polynomial Multiplication (n = 4, q = 17)Negacyclic Wrap
Input Polynomial a
1234
Input Polynomial b
5678
Negacyclic Product a * b mod (X^4 + 1)
121529 Correct Ring Product
Plain Cyclic Wrap (without sign flip)
150159 Wrong (Missing X^n = -1)
Constant Term c[0] Derivation
(1*5 - 2*8 - 3*7 - 4*6) mod 17 = -56 mod 17 = 12

1.3 A lattice cipher in a dozen lines (Toy KEM)

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.

[5]:
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")
Out[5]:
PASSED (0.5 ms)
Toy KEM Keys & Parameters (q = 97, n = 8, eta = 1)Key Generation
Public Matrix a
41195083696812
Public Vector t = a*s + e
8982486392227211
Secret Key s (centred in [-1, +1])
01-11-1-1-10
Toy KEM Transmission [Bit = 0]Bit 0 Transmission
Ephemeral Secret y
-1-1111-111
Injected Noise (e1, e2)
e1=0-1-1-11-100, e2=-1
Ciphertext u (8 ring coefficients)
961456454157228
Ciphertext v (scalar mod 97)
56
Decryption Product s*u[0] mod 97
55
Noisy Message w = v - s*u
1 mod 97 (centred: 1)
Target Center (bit * 48)
0
Decision Boundary (|w| < 24)
|+1| < 24 is True
Recovered Bit
Bit 0 Success: Bit Recovered
Toy KEM Transmission [Bit = 1]Bit 1 Transmission
Ephemeral Secret y
1-11011-1-1
Injected Noise (e1, e2)
e1=111-10-111, e2=-1
Ciphertext u (8 ring coefficients)
638585837931927
Ciphertext v (scalar mod 97)
50
Decryption Product s*u[0] mod 97
4
Noisy Message w = v - s*u
46 mod 97 (centred: 46)
Target Center (bit * 48)
48
Decision Boundary (|w| < 24)
|+46| < 24 is False
Recovered Bit
Bit 1 Success: Bit Recovered

1.4 What Breaks the Scheme: Noise Overflows and Decryption Failures

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.

[6]:
# 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')
Out[6]:
PASSED (141.7 ms)
Concrete Decryption Failure (Trial #4, eta = 4)Noise Overflow
Sent Message Bit
0
Ciphertext u
731933144844929
Ciphertext v
35
Reconstructed Noisy w
30 mod 97 (centred: 30)
Decision Boundary (+/- q/4)
+/- 24
Overflow Check
|+30| >= 24 -> Bit flips from 0 to 1!
Decryption Verdict
Bit Flip Error (Decoding Failure)
Failure Rate vs Noise Bound (1000 trials each, q = 97, n = 8)
Noise Bound (eta)Empirical Failure RateStatus
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%)
[7]:
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")
Out[7]:
PASSED (0.3 ms)
Lattice Security & Brute-Force FeasibilitySearch Space Scaling
Toy Parameters (n = 4, eta = 1)
Search space: 3^4 = 81 candidates
Brute-Force Attack Result
Found secret in 9 attempts Trivially Broken
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 Information-Theoretically Impenetrable

---

2. The Number Theoretic Transform (NTT)

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.

[8]:
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')
Out[8]:
PASSED (26.6 ms)
Roots of Unity and Modulus Factorisation
StandardModulus qPrime Factorisation of (q - 1)256th Root512th Root
ML-KEM33292^8 * 13^1ExistsNone
ML-DSA83804172^13 * 3^1 * 11^1 * 31^1ExistsExists
[9]:
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).")
Out[9]:
PASSED (0.4 ms)
NTT Twiddle Factors & ArchitecturesTwiddle Tables
ML-KEM Generator (order 256)
zeta = 17, n^-1 = 3303
First 8 Bit-Reversed Twiddles
117292580328926426301897848
ML-DSA Generator (order 512)
zeta = 1753, n^-1 = 8347681
First 8 Bit-Reversed Twiddles
14808194376560737615135178923549669152347395178987
ML-KEM stops at 7 stages (128 degree-1 pairs); ML-DSA runs 8 full stages (256 scalar values).
[10]:
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")
Out[10]:
PASSED (1.0 ms)
ML-KEM NTT Verification: INTT(NTT(a)) == aML-KEM Inversion
Input Polynomial a[:12]
720334659728598111124137150
NTT Transformed a_hat[:12]
1199327814572938708749632291031632561736995
Reconstructed INTT(NTT(a))[:12]
720334659728598111124137150
Exact Inversion Check
All 256 coefficients match (Max diff: 0) Verified: Exact Inversion
ML-KEM In-Memory Layout: 128 Pairs of Degree-1 Polynomials
Block NumberMemory IndicesStored Value Pair [a0, a1]Modulus Ideal
Block 0indices ( 0, 1)(1199, 3278)mod (X^2 - 17)
Block 1indices ( 2, 3)(1457, 2938)mod (X^2 - 3312)
Block 2indices ( 4, 5)( 708, 749)mod (X^2 - 2761)
Block 3indices ( 6, 7)( 632, 2910)mod (X^2 - 568)
Block 127indices (254, 255)(2462, 409)mod (X^2 - 1175)
ML-DSA NTT Verification: Full 8-Stage 256-Point InversionML-DSA Inversion
Input Polynomial a[:8]
17148279410541672803934
NTT Transformed a_hat[:8]
63717274612327561007867042081295849541977278194676099772
Reconstructed INTT(NTT(a))[:8]
17148279410541672803934
Exact Inversion Check
All 256 coefficients match (Max diff: 0) Verified: Exact Inversion
[11]:
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')
Out[11]:
PASSED (53.5 ms)
Polynomial Multiplication: Schoolbook vs. Fast NTTArithmetic Equivalence
Input a[:8]
2619456102303711261003914571
Input b[:8]
262138486240922689422410902
Schoolbook Product c[:8]
133120551467538301034317812355
Fast NTT Product c[:8]
133120551467538301034317812355
Equivalence Check
Identical for all 256 coefficients (Max diff: 0) Exact Match
Execution Speedup (Python)
Schoolbook: 3.28 ms | NTT: 0.26 ms -> 12.4x Speedup
Hardware Multiplier Count: Schoolbook vs Fast NTT
Multiplication AlgorithmModular MultipliesTime ComplexityHardware Workload
Schoolbook Convolution256 x 256 = 65,536O(n^2)Baseline (100%)
Fast NTT (ML-KEM)2*896 (fwd) + 640 (ptwise) + 896 (inv) = 3,328O(n log n)95% Workload Reduction
Fast NTT (ML-DSA)2*1024 (fwd) + 256 (ptwise) + 1024 (inv) = 3,328O(n log n)95% Workload Reduction
[12]:
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')
Out[12]:
PASSED (0.2 ms)
8-Point Cooley-Tukey Butterfly Hardware Register Trace
Cooley-Tukey Pipeline StageHardware Register Bank Contents
Stage 0 (Input)12345678
Stage 1 (len=4)15774278424611
Stage 2 (len=2)303102648252323
Stage 3 (Final Output)7978603723737568

2.1 Modular reduction, and why the shape of `q` matters

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.

[13]:
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')
Out[13]:
PASSED (8.3 ms)
Barrett Constant-Multiplier Reduction Validation
StandardModulus qBit Length tPrecomputed Multiplier muBarrett Validation
ML-KEM332912 bitsfloor(2^24 / q) = 5039Worst corrections: 1 10k Passed
ML-DSA838041723 bitsfloor(2^46 / q) = 8396807Worst corrections: 1 10k Passed

---

3. ML-KEM, FIPS 203 <span style="color:#DF8E1D">(Demo 2)</span>

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.

[14]:
# ---- 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.")
Out[14]:
Keccak helpers and samplers ready.
PASSED (0.5 ms)
[15]:
# ---- 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.")
Out[15]:
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)
Polynomial Coefficient Compression & Decompression (Concrete Numbers)Lossy Compression
Original Polynomial p[:8] (mod 3329)
782127725584001825132512782793
Compressed to d=4 bits (values in [0, 15])
4612296613
Decompressed from d=4 bits
832124824974161873124812482705
Compression Errors |p - Decomp(Comp(p))|[:8]
5029611648773088
Max Error Observed (d=4)
104 levels (Theoretical Bound: ceil(q/2^5) = 105) Small Error Tolerated
Compressed to d=10 bits (values in [0, 1023])
241393787123561408393859
Decompressed from d=10 bits
783127825594001824132612782793
Max Error Observed (d=10)
2 levels (Theoretical Bound: ceil(q/2^11) = 2) Near-Exact Recovery
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.
[16]:
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.")
Out[16]:
ML-KEM implemented.
PASSED (0.8 ms)
[17]:
# ---- 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()
Out[17]:
  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)
[18]:
# ---- 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")
Out[18]:
==============================================================================
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)
ML-KEM-768 Key Generation SummaryKeyGen Stage
Seed rho (32 bytes)
a75da6ea6bd3360b...
Secret Vector s[0][:8]
00-1-1-1-12-1
Noise Vector e[0][:8]
1-11101-10
Public Key t-hat[0][:8]
22601352310679269422532095128
Published ek Size
1184 bytes (t-hat: 1152 B + rho: 32 B)
ML-KEM-768 Encapsulation SummaryEncaps Stage
Plaintext Message m
f87db3d355f5b90b...
Ephemeral Vector y[0][:8]
-1-10-10201
Ciphertext Vector u[0][:8]
132113272503933327303317542216
Ciphertext Scalar v[:8]
98519742346132224616785741121
Compressed Ciphertext Size
1088 bytes (u: 960 B, v: 128 B)
Bob Shared Secret K_bob
caf076540ad507f454f3d338...
ML-KEM-768 Decapsulation & AgreementDecaps Stage
Observed Max Noise
211 (Bound q/4 = 832)
Safety Headroom Remaining
621 levels
Decoded Message Match
m_rec == m: True Message Recovered
Alice Shared Secret K_alice
caf076540ad507f454f3d338...
Key Synchronization Match
K_alice == K_bob: True Keys Synchronized
[19]:
# ---- 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')
Out[19]:
PASSED (3976.9 ms)
Noise Injection Attack: Perturbing Ciphertext (+1665 mod 3329)FO Defense
Original Message Bit 0
1
Decoded Bit 0 (clean ct)
1 Correct
Decoded Bit 0 (faulted ct)
0 Bit Flipped
Bob True Shared Key
ae4a78cdd9151255e62f25b3...
Alice Decapsulated Key
e3223209ed3dc021827c66c2...
Key Synchronization Result
Match? False Silent Desynchronization (FO Trapdoor)
ML-KEM Noise Sweep & Fujisaki-Okamoto Defense (60 trials each)
Injected Error (delta)Message Bit Flip RateFO Implicit RejectionObserved 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 q100.0 %100.0 %Noise Overflow (100% Flip)
+1665 mod q100.0 %100.0 %Noise Overflow (100% Flip)
[20]:
# ---- 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.")
Out[20]:
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)
[21]:
# ---- 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().")
Out[21]:
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)

---

4. ML-DSA, FIPS 204 <span style="color:#DF8E1D">(Demo 3)</span>

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.

[22]:
# ---- 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.")
Out[22]:
Power2Round, Decompose and the MakeHint/UseHint identity: verified on 3000 values.
PASSED (8.2 ms)
[23]:
# ---- 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.")
Out[23]:
ML-DSA rounding, packing and samplers ready.
PASSED (0.3 ms)
[24]:
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.")
Out[24]:
ML-DSA implemented.
PASSED (1.5 ms)
[25]:
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()
Out[25]:
  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)
[26]:
# ---- 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")
Out[26]:
==============================================================================
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)
ML-DSA-65 Key Generation SummaryKeyGen Stage
Seed rho (32 bytes)
420ed13318e20bc8...
Secret Vector s1[0][:8]
32-1-442-3-3
Secret Vector s2[0][:8]
-232-13-21-4
Verification Key pk Size
1952 bytes (t1 packed + rho)
Signing Key sk Size
4032 bytes (rho + key + tr + s1 + s2 + t0)
ML-DSA-65 Signature Execution (Accepted on Attempt #5)Signature Generated
Signing Attempts Required
5 attempt(s)
Challenge Hash c_tilde
d9428142022da25d1076ac60...
Signature Vector z[0][:8]
66823519715484182105260208959284616-353888-52787
Norm Bound ||z||_∞
523,977 < 524,092 Norm Satisfied
Hint Bits Set in Matrix h
36 / 55 Hints Valid
Encoded Signature Size
3309 bytes (c_tilde: 48B, z: 3200B, h: 61B)
ML-DSA-65 Verification: 4 Architectural Checks (Slide 64 / FIPS 204)4-Condition Verification
Condition 1: Response Vector Norm
||z||_∞ = 523,977 < 524,092 Condition 1 Pass
Condition 2: Hint Matrix Weight
wt(h) = 36 <= 55 (bound omega) Condition 2 Pass
Condition 3: High-Bits Recovery
w'1 == w1 in all 6 polynomials Condition 3 Pass
Condition 4: Challenge Hash Match
c' == c~ (d9428142022da25d... == d9428142022da25d...) Condition 4 Pass
Overall Verification Verdict
All 4 Conditions Evaluated True -> ACCEPT (⊤) Cryptographically Valid
Reference: Slide 64 (ML-DSA Sign Part 2) & FIPS 204 Algorithm 3. Verification requires all 4 conditions to evaluate true.
ML-DSA-65 Tamper Resistance TestsTamper Tests
Genuine Firmware Image
True Accept Signature
One Byte Flipped
False Reject Signature (Tampered)
Wrong Execution Context (ctx=b'boot')
False Reject Signature (Wrong Context)
[27]:
# ---- 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.")
Out[27]:
PASSED (53.2 ms)
ML-DSA-65 Signing Attempt #1 [REJECTED]Attempt #1 Aborted
Input Message
'FIRMWARE v3.2.0-PATCH-2026'
Attempt Verdict
Rejected (Retry Next Candidate)
Rejection Reason
||z||_∞ = 524,222 >= bound 524,092 (would leak secret s1 bits)
Candidate Mask y / z[0][:6]
-17795372840416909302129109747444789
Challenge Hash c_tilde
d8e3f5d332e2e52e67fdfd15...
Security Rationale
If this signature were published, an adversary collecting samples could compute the private key s1/s2.
Fiat-Shamir with Aborts: Discard candidate and re-sample mask vector y without incrementing any public state.
ML-DSA-65 Signing Attempt #2 [REJECTED]Attempt #2 Aborted
Input Message
'FIRMWARE v3.2.0-PATCH-2026'
Attempt Verdict
Rejected (Retry Next Candidate)
Rejection Reason
||z||_∞ = 524,176 >= bound 524,092 (would leak secret s1 bits)
Candidate Mask y / z[0][:6]
-311646286737-200994-43876675079-191220
Challenge Hash c_tilde
5043b976354df1bda77cf2a3...
Security Rationale
If this signature were published, an adversary collecting samples could compute the private key s1/s2.
Fiat-Shamir with Aborts: Discard candidate and re-sample mask vector y without incrementing any public state.
ML-DSA-65 Signing Attempt #3 [REJECTED]Attempt #3 Aborted
Input Message
'FIRMWARE v3.2.0-PATCH-2026'
Attempt Verdict
Rejected (Retry Next Candidate)
Rejection Reason
||z||_∞ = 524,283 >= bound 524,092 (would leak secret s1 bits)
Candidate Mask y / z[0][:6]
-485318-337368-433470-421480-273511-434411
Challenge Hash c_tilde
038ab212fdee7cbcbb038a70...
Security Rationale
If this signature were published, an adversary collecting samples could compute the private key s1/s2.
Fiat-Shamir with Aborts: Discard candidate and re-sample mask vector y without incrementing any public state.
ML-DSA-65 Signing Attempt #4 [REJECTED]Attempt #4 Aborted
Input Message
'FIRMWARE v3.2.0-PATCH-2026'
Attempt Verdict
Rejected (Retry Next Candidate)
Rejection Reason
||z||_∞ = 524,201 >= bound 524,092 (would leak secret s1 bits)
Candidate Mask y / z[0][:6]
-35875-340475-28056244140-98951518360
Challenge Hash c_tilde
f5d19319f3221a9354635db0...
Security Rationale
If this signature were published, an adversary collecting samples could compute the private key s1/s2.
Fiat-Shamir with Aborts: Discard candidate and re-sample mask vector y without incrementing any public state.
ML-DSA-65 Signing Attempt #5 [REJECTED]Attempt #5 Aborted
Input Message
'FIRMWARE v3.2.0-PATCH-2026'
Attempt Verdict
Rejected (Retry Next Candidate)
Rejection Reason
||r0||_∞ = 261,736 >= bound 261,692 (low bits would fail high-bit hint recovery)
Candidate Mask y / z[0][:6]
-10750-410509391086-333761107784-8067
Challenge Hash c_tilde
e357778295990c0f87ea247f...
Security Rationale
If this signature were published, an adversary collecting samples could compute the private key s1/s2.
Fiat-Shamir with Aborts: Discard candidate and re-sample mask vector y without incrementing any public state.
ML-DSA-65 Signing Attempt #6 [REJECTED]Attempt #6 Aborted
Input Message
'FIRMWARE v3.2.0-PATCH-2026'
Attempt Verdict
Rejected (Retry Next Candidate)
Rejection Reason
||r0||_∞ = 261,882 >= bound 261,692 (low bits would fail high-bit hint recovery)
Candidate Mask y / z[0][:6]
351792159839-173330-49998139471448527
Challenge Hash c_tilde
28e55267aff3929841f65a2f...
Security Rationale
If this signature were published, an adversary collecting samples could compute the private key s1/s2.
Fiat-Shamir with Aborts: Discard candidate and re-sample mask vector y without incrementing any public state.
ML-DSA-65 Signing Attempt #7 [REJECTED]Attempt #7 Aborted
Input Message
'FIRMWARE v3.2.0-PATCH-2026'
Attempt Verdict
Rejected (Retry Next Candidate)
Rejection Reason
||z||_∞ = 524,108 >= bound 524,092 (would leak secret s1 bits)
Candidate Mask y / z[0][:6]
-78409-42069-2266078960505867-7966
Challenge Hash c_tilde
19e78bc238dee9e65e7855b5...
Security Rationale
If this signature were published, an adversary collecting samples could compute the private key s1/s2.
Fiat-Shamir with Aborts: Discard candidate and re-sample mask vector y without incrementing any public state.
ML-DSA-65 Signing Attempt #8 [ACCEPTED & VERIFIED]Valid Signature (Attempt #8)
Input Message
'FIRMWARE v3.2.0-PATCH-2026'
Attempt Verdict
Accepted on Attempt #8!
Norm Check ||z||_∞
524,068 < 524,092 Strictly Bounded
Low-Bits Check ||r0||_∞
261,491 < 261,692 Hints Reliable
Valid Hints Set in h
30 / 55 Hints Fit Packing
Final Signature Vector z[0][:6]
305862267306346844495509403296-322369
Verification Check
verify(pk, msg, sig) == True Cryptographically Valid
Total loop iterations required: 8. Secret key privacy mathematically preserved.
[28]:
# ---- 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')
Out[28]:
PASSED (3280.4 ms)
ML-DSA Rejection Loop Dynamics across Parameter Sets
Security LevelMean AttemptsMedianMin / Maxz Rejection %r0 Rejection %Hint Rejection %
ML-DSA-444.10 (spec: 4.25)3.01 / 1658.1 %41.1 % 0.8 %
ML-DSA-654.72 (spec: 5.10)4.01 / 1555.0 %45.0 % 0.0 %
ML-DSA-873.85 (spec: 3.85)3.01 / 1448.2 %50.9 % 0.9 %

---

5. Profiling & Where the Cycles Go

Now that Keccak and NTT are instrumented, convert the operation counts into estimated hardware cycles.

[29]:
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.")
Out[29]:
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)
[30]:
# 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.")
Out[30]:
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)
Hardware Cycle Split Estimate (1 Butterfly Engine + 1-Round/Cycle Keccak)
OperationKeccak (cyc)NTT (cyc)Total (cyc)Keccak %NTT %
ML-KEM-768 KeyGen1,0325,3766,408 16.1 % 83.9 %
ML-KEM-768 Encaps1,0566,2727,328 14.4 % 85.6 %
ML-KEM-768 Decaps1,2729,85611,128 11.4 % 88.6 %
ML-DSA-65 KeyGen5,76011,26417,024 33.8 % 66.2 %
ML-DSA-65 Sign7,824117,760125,584 6.2 % 93.8 %
ML-DSA-65 Verify4,96818,43223,400 21.2 % 78.8 %

---

6. Hardware Cost Estimator

Estimate latency and throughput based on butterfly parallelism and Keccak rounds per cycle.

[31]:
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)
Out[31]:
PASSED (0.2 ms)
Parallel Butterfly Scaling for ML-KEM-768 Encapsulation
Butterfly EnginesOnly NTT ScaledBoth Scaled (Keccak + NTT)Latency @ 100 MHz
1 units7,328 cyc7,328 cyc73.3 us
2 units2,795 cyc2,267 cyc22.7 us
4 units1,538 cyc746 cyc7.5 us
8 units1,189 cyc265 cyc2.6 us
16 units1,093 cyc103 cyc1.0 us
32 units1,074 cyc62 cyc0.6 us
[32]:
# 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')
Out[32]:
PASSED (1.8 ms)
Representative Hardware Architecture Profiles (ML-KEM-768 Encaps)
Design Point ProfileHardware ResourcesClockTotal CyclesLatency
Tiny (Resource-Constrained IoT)1 BF / 0.5 rd/cyc50 MHz8,3840.17 ms
Balanced (Embedded SoC)2 BF / 1 rd/cyc100 MHz4,5350.05 ms
Fast (High-Throughput HSM)8 BF / 2 rd/cyc300 MHz1,5980.01 ms

---

7. Side Channels: Leaky vs Constant-Time Comparison

The Fujisaki-Okamoto transform makes ML-KEM IND-CCA2 secure, provided the ciphertext comparison inside decapsulation is constant-time.

[33]:
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)
Out[33]:
PASSED (147.6 ms)
Ciphertext Comparison Timing: Early-Exit Leak vs Constant-Time
Matching Prefix LengthLeaky Compare TimeSafe 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
[34]:
# 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")
Out[34]:
PASSED (1.5 ms)
Timing Side-Channel Exploitation: Key RecoveryLive Timing Exploit
Target Secret (16 bytes)
0xf800d5f44e39621b6f67bb5ec8d4e347
Recovered Secret via Timing
0xf800d5f44e39621b6f67bb5ec8d4e347 Exact Recovery
Total Oracle Timing Queries
2186 queries
Brute Force Search Space
2^128 ≈ 3.4 x 10^38 attempts
Attack Complexity Reduction
Reduced from 2^128 to ~2186 queries Trivially Recovered in 2186 Steps
Core Takeaway
Early-exit ciphertext comparisons completely bypass IND-CCA2 security. Decaps must always execute in constant time.
[35]:
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")
Out[35]:
PASSED (373.1 ms)
ML-DSA Signing Timing Spread (Variable Loop Latency)Signing Latency Spread
Sample Message Count
20 messages
Attempt Distribution
Min: 1 | Median: 3.0 | Max: 19
Execution Time Spread
Min: 6.0 ms | Median: 11.6 ms | Max: 58.2 ms
Timing Channel Implication
Signing latency inherently varies due to rejection sampling; constant-time loops require dummy iterations.

---

8. Interactive Experiments & Things to Break

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.