134 lines
3.4 KiB
Python
134 lines
3.4 KiB
Python
"""
|
|
lib/crypto_utils.py — common CTF crypto helpers (RSA / lattice / misc).
|
|
|
|
Patterns distilled from p4-team/ctf writeups.
|
|
"""
|
|
import math
|
|
from math import gcd, isqrt
|
|
|
|
|
|
def egcd(a, b):
|
|
if b == 0:
|
|
return (a, 1, 0)
|
|
g, x, y = egcd(b, a % b)
|
|
return (g, y, x - (a // b) * y)
|
|
|
|
|
|
def modinv(a, m):
|
|
g, x, _ = egcd(a % m, m)
|
|
if g != 1:
|
|
raise ValueError("modinv: no inverse")
|
|
return x % m
|
|
|
|
|
|
def isqrt(n):
|
|
return math.isqrt(n)
|
|
|
|
|
|
def factor_trivial(n):
|
|
"""Tiny factor finder for small/weak moduli."""
|
|
for p in range(2, 1 << 20):
|
|
if n % p == 0:
|
|
return p, n // p
|
|
return None
|
|
|
|
|
|
# ---- RSA recovery recipes (from p4 writeups) ----
|
|
|
|
def recover_n_from_keys(e, d, ipmq, iqmp):
|
|
"""
|
|
From p4 'lost_modulus': we know e, d, ipmq=modinv(p,q), iqmp=modinv(q,p)
|
|
but NOT n. Recover n via quadratic equation on phi. Returns (p, q) or None.
|
|
"""
|
|
try:
|
|
import gmpy2
|
|
except Exception:
|
|
raise SystemExit("gmpy2 required for recover_n_from_keys")
|
|
|
|
def find_phi(e, d):
|
|
kfi = e * d - 1
|
|
k = kfi // (int(d) * 3)
|
|
while True:
|
|
fi = kfi // k
|
|
try:
|
|
d0 = gmpy2.invert(e, fi)
|
|
if d == d0:
|
|
yield fi
|
|
except Exception:
|
|
pass
|
|
k += 1
|
|
|
|
def solve(ipmq, iqmp, possible_phi):
|
|
a = iqmp - 1
|
|
b = ipmq + iqmp - 2 - possible_phi
|
|
c = ipmq * possible_phi - possible_phi
|
|
delta = b * b - 4 * a * c
|
|
if delta > 0:
|
|
r, correct = gmpy2.iroot(delta, 2)
|
|
if correct:
|
|
for x in [(-b - r) // (2 * a), (-b + r) // (2 * a)]:
|
|
if gmpy2.is_prime(x + 1):
|
|
q = x + 1
|
|
p = possible_phi // x + 1
|
|
return int(p), int(q)
|
|
return None
|
|
|
|
for phi in find_phi(e, d):
|
|
res = solve(ipmq, iqmp, phi)
|
|
if res:
|
|
return res
|
|
return None
|
|
|
|
|
|
def common_modulus_attack(c1, c2, e1, e2, n):
|
|
"""Same message encrypted with same n, coprime exponents."""
|
|
g, a, b = egcd(e1, e2)
|
|
if g != 1:
|
|
raise ValueError("e1,e2 not coprime")
|
|
if a < 0:
|
|
c1, a = modinv(c1, n), -a
|
|
if b < 0:
|
|
c2, b = modinv(c2, n), -b
|
|
m = (pow(c1, a, n) * pow(c2, b, n)) % n
|
|
return m
|
|
|
|
|
|
def hastad_broadcast(cts, es, n, mlen=1):
|
|
"""CRT-combine same small message raised to small exponents e across moduli.
|
|
cts[k] = m^es[k] mod n[k]. Returns m if m^max(e) < n_prod."""
|
|
from functools import reduce
|
|
N = reduce(lambda a, b: a * b, n)
|
|
result = 0
|
|
for c, ni in zip(cts, n):
|
|
Ni = N // ni
|
|
result = (result + c * Ni * modinv(Ni, ni)) % N
|
|
k = max(es)
|
|
return int(round(result ** (1.0 / k)))
|
|
|
|
|
|
def wiener(e, n):
|
|
"""Wiener's attack: small d. Returns d or None."""
|
|
def cf(a, b):
|
|
while b:
|
|
yield a // b
|
|
a, b = b, a % b
|
|
def convergents(cf_gen):
|
|
h0, h1 = 0, 1
|
|
k0, k1 = 1, 0
|
|
for q in cf_gen:
|
|
h0, h1 = h1, q * h1 + h0
|
|
k0, k1 = k1, q * k1 + k0
|
|
yield h1, k1
|
|
for k, d in convergents(cf(e, n)):
|
|
if k == 0:
|
|
continue
|
|
if (e * d - 1) % k == 0:
|
|
phi = (e * d - 1) // k
|
|
s = n - phi + 1
|
|
disc = s * s - 4 * n
|
|
if disc >= 0:
|
|
r = isqrt(disc)
|
|
if r * r == disc and (s + r) % 2 == 0:
|
|
return d
|
|
return None
|