""" 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