Files
TOIFAS/code/check_cov_attack.py
T
KiHoLee 3529ab1918 Audit round: fair OMA reference, dense grids, covariance-attack checks
Resource-match the OMA reference in the key-length sweep (oma_ser_keylen),
which gives it the L/16 combining gain the longer frame allows. The
proposal now passes a resource-matched OMA by 1.27x at L=64 rather than
the 4.3x reported against a fixed-d reference.

Densify the JSR, sensitivity, and brute-force grids so the curves are
smooth, give the index cipher its channel floor instead of error-free
reception, and add the permutation-key known-plaintext attack
(exp_permkpa) so Fig. 7 carries a conventional linear scheme.

Add check_cov_attack.py and check_cov_ceiling.py: a referee raised a
ciphertext-only second-order attack; the exact-population test shows the
received covariance leaks only a sparse rank-deficient subset of the key
Gram and leaves the eavesdropper at the random-guess level.

Dump verify_math.csv, move the superseded V=256 pilot CSVs to data/pilot.
2026-08-16 23:48:39 +09:00

113 lines
4.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Independent check of the ciphertext-only second-order attack (audit M1).
Claim under test: an eavesdropper who observes only received frames (no
known indices) can estimate the key Gram matrix M^T M from the sample
covariance, because per period
E[y_k y_l] = (1/c^2) * E[h^2] * C_kl * (M^T M)_kl,
where C_kl = (1/Vu) sum_i b_{i,k} b_{i,l} is the PUBLIC codebook column
correlation and the noise touches only the diagonal.
Procedure, using nothing the threat model keeps secret:
1. collect N received frames y_n = h_n * (1/c) sum_u e_{s_u} ⊙ m_u + noise
2. form the per-period sample second moment S_kl = mean_n y_{n,k} y_{n,l}
3. divide the off-diagonal by C_kl (public) to get G_hat ≈ M^T M
4. set the diagonal of G_hat to U (unit-modulus keys)
5. factor G_hat = M_hat^T M_hat (rank U), then for Walsh-Hadamard keys
round to ±1 and search the 2^U U! signed permutations, keeping the
M_hat that best decodes a handful of the collected frames
6. report the recovered-entry fraction and the eavesdropper SER, both
WITHOUT ever using a known index
Run under WSL. Prints a verdict; writes nothing to data/.
"""
from __future__ import annotations
import itertools
import math
import numpy as np
import torch
from sse_lib import rayleigh_gain, DEVICE
from exp_full import get_model, hadamard, eval_ser_eve
def collect_frames(m, n, snr_db, seed):
"""Received frames and the true indices (indices kept only for scoring)."""
g = torch.Generator().manual_seed(seed)
Bn = m.unit_codebook()
true_m = m.masks()
c = m.c
sigma = math.sqrt(1.0 / (m.d * 10.0 ** (snr_db / 10.0)))
digits = torch.randint(m.vu, (n, m.users, m.P), generator=g).to(DEVICE)
e = Bn[digits] / math.sqrt(m.P) # (n,U,P,L)
y = (e * true_m[None, :, None, :]).sum(dim=1) / c # (n,P,L)
h = rayleigh_gain((n,), device=DEVICE)
y = h[:, None, None] * y + sigma * torch.randn(n, m.P, m.L, device=DEVICE)
return y, digits, Bn, true_m
def codebook_corr(Bn):
"""Public column correlation C_kl = (1/Vu) sum_i b_ik b_il."""
return (Bn.T @ Bn) / Bn.shape[0] # (L,L)
def attack(m, snr_db, n_frames, seed):
y, digits, Bn, true_m = collect_frames(m, n_frames, snr_db, seed)
U, L = m.users, m.L
yf = y.reshape(-1, L) # pool all periods
S = (yf.T @ yf) / yf.shape[0] # (L,L) 2nd moment
C = codebook_corr(Bn) # public
G = torch.zeros(L, L, device=DEVICE)
mask = C.abs() > 1e-3
G[mask] = S[mask] / C[mask] # ≈ (1/c^2) M^T M
scale = float(torch.diagonal(G)[mask.diagonal()].mean()) / U
G = G / max(scale, 1e-9) # normalize so diag≈U
G.fill_diagonal_(float(U)) # unit-modulus keys
# symmetric rank-U factor
G = 0.5 * (G + G.T)
evals, evecs = torch.linalg.eigh(G)
idx = torch.argsort(evals, descending=True)[:U]
root = evecs[:, idx] * evals[idx].clamp_min(0).sqrt()
Mhat0 = root.T # (U,L), up to U×U orth
# for WH keys, snap to ±1 and search signed row permutations
cand = torch.sign(Mhat0)
cand[cand == 0] = 1.0
best = None
best_ser = 1.0
val = torch.arange(min(2000, n_frames))
for perm in itertools.permutations(range(U)):
for signs in itertools.product([1.0, -1.0], repeat=U):
Mh = (cand[list(perm)] *
torch.tensor(signs, device=DEVICE)[:, None])
ser = eval_ser_eve(m, Mh.cpu(), [snr_db], frames=20_000,
seed=13)[0]
if ser < best_ser:
best_ser, best = ser, Mh
# recovered-entry fraction against the true keys (best sign-aligned)
tm = torch.sign(true_m).to(DEVICE)
frac = 0.0
for perm in itertools.permutations(range(U)):
for signs in itertools.product([1.0, -1.0], repeat=U):
Mh = (best[list(perm)] *
torch.tensor(signs, device=DEVICE)[:, None])
frac = max(frac, float((Mh == tm).float().mean()))
return frac, best_ser
def main():
U, L = 4, 16
K0 = torch.tensor(hadamard(L)[1:U + 1], dtype=torch.float32)
m = get_model(iters=4000, freeze_W=K0)
m.eval()
chance = 1.0 - (1.0 / m.vu) ** m.P
print(f"chance SER = {chance:.5f}, legitimate reference ~0.276")
print("ciphertext-only (NO known plaintext):")
for snr in (10.0, 20.0):
for nf in (300, 1000, 10000):
frac, ser = attack(m, snr, nf, seed=1234 + nf)
print(f" {snr:4.0f} dB N={nf:6d} "
f"key-entry recovery={frac:.3f} eve SER={ser:.4f}")
if __name__ == "__main__":
main()