crypto-gu 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
crypto_gu/__init__.py ADDED
@@ -0,0 +1,12 @@
1
+ """crypto_gu: a pure standard-library cryptography toolkit.
2
+
3
+ Research- and teaching-oriented primitives (hashes, KDFs, ciphers, AEAD,
4
+ RSA padding schemes) alongside deterministic attack implementations
5
+ (padding oracles, small-exponent and factorisation helpers). No third
6
+ party dependencies; every primitive is validated against RFC/NIST test
7
+ vectors.
8
+ """
9
+
10
+ __version__ = "0.1.0"
11
+
12
+ __all__ = ["__version__"]
File without changes
@@ -0,0 +1,212 @@
1
+ """RSA encryption and signature schemes from PKCS #1 v2.2 (RFC 8017).
2
+
3
+ Provides MGF1, RSAES-OAEP (section 7.1) and RSASSA-PSS (section 8.1). The
4
+ underlying trapdoor is the textbook :class:`~crypto_gu.asymmetric.rsa.RSAKey`
5
+ already in this package; these helpers add the randomised, provably related
6
+ padding that turns it into a real-world scheme.
7
+
8
+ Pure Python, standard library only (``os.urandom`` supplies the randomness).
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import os
14
+
15
+ from crypto_gu.errors import DecryptionError, InvalidSignatureError
16
+ from crypto_gu.hashes import HASH_TABLE
17
+
18
+
19
+ def _hash_length(hash_name: str) -> int:
20
+ key = hash_name.lower()
21
+ if key not in HASH_TABLE:
22
+ raise ValueError("unknown hash algorithm: %r" % hash_name)
23
+ return len(HASH_TABLE[key](b""))
24
+
25
+
26
+ def _digest(hash_name: str, data: bytes) -> bytes:
27
+ return HASH_TABLE[hash_name.lower()](data)
28
+
29
+
30
+ def _i2osp(value: int, length: int) -> bytes:
31
+ if value < 0 or value >> (8 * length):
32
+ raise ValueError("integer too large")
33
+ return value.to_bytes(length, "big")
34
+
35
+
36
+ def _os2ip(data: bytes) -> int:
37
+ return int.from_bytes(data, "big")
38
+
39
+
40
+ def mgf1(seed: bytes, length: int, hash_name: str = "sha256") -> bytes:
41
+ """Mask generation function MGF1 (RFC 8017 appendix B.2.1)."""
42
+ if length < 0:
43
+ raise ValueError("mask length must be non-negative")
44
+ hlen = _hash_length(hash_name)
45
+ if length > (1 << 32) * hlen:
46
+ raise ValueError("mask too long")
47
+ output = b""
48
+ counter = 0
49
+ while len(output) < length:
50
+ output += _digest(hash_name, seed + counter.to_bytes(4, "big"))
51
+ counter += 1
52
+ return output[:length]
53
+
54
+
55
+ def _xor(left: bytes, right: bytes) -> bytes:
56
+ return bytes(a ^ b for a, b in zip(left, right))
57
+
58
+
59
+ def _modulus_size(n: int) -> int:
60
+ return (n.bit_length() + 7) // 8
61
+
62
+
63
+ # --------------------------------------------------------------------------- #
64
+ # RSAES-OAEP
65
+ # --------------------------------------------------------------------------- #
66
+
67
+
68
+ def oaep_encode(message: bytes, k: int, hash_name: str = "sha256", label: bytes = b"") -> bytes:
69
+ """EME-OAEP encoding (RFC 8017 section 7.1.1)."""
70
+ hlen = _hash_length(hash_name)
71
+ if len(message) > k - 2 * hlen - 2:
72
+ raise ValueError("message too long for OAEP")
73
+ lhash = _digest(hash_name, label)
74
+ ps = b"\x00" * (k - len(message) - 2 * hlen - 2)
75
+ db = lhash + ps + b"\x01" + message
76
+ seed = os.urandom(hlen)
77
+ db_mask = mgf1(seed, k - hlen - 1, hash_name)
78
+ masked_db = _xor(db, db_mask)
79
+ seed_mask = mgf1(masked_db, hlen, hash_name)
80
+ masked_seed = _xor(seed, seed_mask)
81
+ return b"\x00" + masked_seed + masked_db
82
+
83
+
84
+ def oaep_decode(encoded: bytes, k: int, hash_name: str = "sha256", label: bytes = b"") -> bytes:
85
+ """EME-OAEP decoding; raises :class:`DecryptionError` on malformed input."""
86
+ hlen = _hash_length(hash_name)
87
+ if len(encoded) != k or k < 2 * hlen + 2:
88
+ raise DecryptionError("decryption error")
89
+ if encoded[0] != 0:
90
+ raise DecryptionError("decryption error")
91
+ masked_seed = encoded[1 : 1 + hlen]
92
+ masked_db = encoded[1 + hlen :]
93
+ seed = _xor(masked_seed, mgf1(masked_db, hlen, hash_name))
94
+ db = _xor(masked_db, mgf1(seed, k - hlen - 1, hash_name))
95
+ if db[:hlen] != _digest(hash_name, label):
96
+ raise DecryptionError("decryption error")
97
+ index = hlen
98
+ while index < len(db) and db[index] == 0:
99
+ index += 1
100
+ if index >= len(db) or db[index] != 1:
101
+ raise DecryptionError("decryption error")
102
+ return db[index + 1 :]
103
+
104
+
105
+ def oaep_encrypt(key, message: bytes, hash_name: str = "sha256", label: bytes = b"") -> bytes:
106
+ """RSAES-OAEP encrypt with a public (or full) :class:`RSAKey`."""
107
+ k = _modulus_size(key.n)
108
+ encoded = oaep_encode(message, k, hash_name, label)
109
+ m = _os2ip(encoded)
110
+ if m >= key.n:
111
+ raise ValueError("encoded message representative out of range")
112
+ return _i2osp(pow(m, key.e, key.n), k)
113
+
114
+
115
+ def oaep_decrypt(key, ciphertext: bytes, hash_name: str = "sha256", label: bytes = b"") -> bytes:
116
+ """RSAES-OAEP decrypt; requires the private exponent."""
117
+ if key.d is None:
118
+ raise DecryptionError("decryption error")
119
+ k = _modulus_size(key.n)
120
+ if len(ciphertext) != k:
121
+ raise DecryptionError("decryption error")
122
+ c = _os2ip(ciphertext)
123
+ if c >= key.n:
124
+ raise DecryptionError("decryption error")
125
+ m = pow(c, key.d, key.n)
126
+ return oaep_decode(_i2osp(m, k), k, hash_name, label)
127
+
128
+
129
+ # --------------------------------------------------------------------------- #
130
+ # RSASSA-PSS
131
+ # --------------------------------------------------------------------------- #
132
+
133
+
134
+ def pss_encode(message: bytes, em_bits: int, hash_name: str = "sha256", salt_length: int = 32) -> bytes:
135
+ """EMSA-PSS encoding (RFC 8017 section 9.1.1)."""
136
+ hlen = _hash_length(hash_name)
137
+ em_len = (em_bits + 7) // 8
138
+ if em_len < hlen + salt_length + 2:
139
+ raise ValueError("encoding error")
140
+ m_hash = _digest(hash_name, message)
141
+ salt = os.urandom(salt_length)
142
+ h = _digest(hash_name, b"\x00" * 8 + m_hash + salt)
143
+ ps = b"\x00" * (em_len - salt_length - hlen - 2)
144
+ db = ps + b"\x01" + salt
145
+ db_mask = mgf1(h, em_len - hlen - 1, hash_name)
146
+ masked_db = bytearray(_xor(db, db_mask))
147
+ masked_db[0] &= 0xFF >> (8 * em_len - em_bits)
148
+ return bytes(masked_db) + h + b"\xbc"
149
+
150
+
151
+ def pss_verify(message: bytes, encoded: bytes, em_bits: int, hash_name: str = "sha256",
152
+ salt_length: int = 32) -> bool:
153
+ """EMSA-PSS verification (RFC 8017 section 9.1.2). Returns a bool."""
154
+ hlen = _hash_length(hash_name)
155
+ em_len = (em_bits + 7) // 8
156
+ if len(encoded) != em_len or em_len < hlen + salt_length + 2:
157
+ return False
158
+ if encoded[-1] != 0xBC:
159
+ return False
160
+ masked_db = bytearray(encoded[: em_len - hlen - 1])
161
+ h = encoded[em_len - hlen - 1 : em_len - 1]
162
+ if masked_db[0] & ~(0xFF >> (8 * em_len - em_bits)):
163
+ return False
164
+ masked_db[0] &= 0xFF >> (8 * em_len - em_bits)
165
+ db = bytearray(_xor(bytes(masked_db), mgf1(h, em_len - hlen - 1, hash_name)))
166
+ db[0] &= 0xFF >> (8 * em_len - em_bits)
167
+ ps_len = em_len - hlen - salt_length - 2
168
+ if db[:ps_len] != b"\x00" * ps_len or db[ps_len] != 0x01:
169
+ return False
170
+ salt = bytes(db[-salt_length:]) if salt_length else b""
171
+ m_hash = _digest(hash_name, message)
172
+ return h == _digest(hash_name, b"\x00" * 8 + m_hash + salt)
173
+
174
+
175
+ def pss_sign(key, message: bytes, hash_name: str = "sha256", salt_length: int = 32) -> bytes:
176
+ """RSASSA-PSS sign; requires the private exponent."""
177
+ if key.d is None:
178
+ raise InvalidSignatureError("signing requires a private key")
179
+ k = _modulus_size(key.n)
180
+ em_bits = key.n.bit_length() - 1
181
+ encoded = pss_encode(message, em_bits, hash_name, salt_length)
182
+ m = _os2ip(encoded)
183
+ if m >= key.n:
184
+ raise ValueError("encoded message representative out of range")
185
+ return _i2osp(pow(m, key.d, key.n), k)
186
+
187
+
188
+ def pss_verify_signature(key, message: bytes, signature: bytes, hash_name: str = "sha256",
189
+ salt_length: int = 32) -> bool:
190
+ """RSASSA-PSS verify with a public (or full) :class:`RSAKey`."""
191
+ k = _modulus_size(key.n)
192
+ if len(signature) != k:
193
+ return False
194
+ s = _os2ip(signature)
195
+ if s >= key.n:
196
+ return False
197
+ m = pow(s, key.e, key.n)
198
+ em_bits = key.n.bit_length() - 1
199
+ return pss_verify(message, _i2osp(m, k), em_bits, hash_name, salt_length)
200
+
201
+
202
+ __all__ = [
203
+ "mgf1",
204
+ "oaep_encode",
205
+ "oaep_decode",
206
+ "oaep_encrypt",
207
+ "oaep_decrypt",
208
+ "pss_encode",
209
+ "pss_verify",
210
+ "pss_sign",
211
+ "pss_verify_signature",
212
+ ]
@@ -0,0 +1,108 @@
1
+ """RSA public-key primitives.
2
+
3
+ Pure Python, standard library only. Key generation uses the Miller-Rabin
4
+ based :mod:`crypto_gu.number_theory` primes. Textbook RSA helpers are
5
+ provided, plus :func:`construct_private` to build a private key from an
6
+ explicit factorisation (the output of the attacks in
7
+ :mod:`crypto_gu.attacks.rsa`).
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from dataclasses import dataclass
13
+
14
+ from crypto_gu import number_theory as nt
15
+
16
+
17
+ def _i2b(value: int, size: int) -> bytes:
18
+ return value.to_bytes(size, "big")
19
+
20
+
21
+ def _b2i(data: bytes) -> int:
22
+ return int.from_bytes(data, "big")
23
+
24
+
25
+ @dataclass
26
+ class RSAKey:
27
+ """A minimal RSA public/private key pair. ``d`` may be ``None``."""
28
+
29
+ n: int
30
+ e: int = 65537
31
+ d: int | None = None
32
+ p: int | None = None
33
+ q: int | None = None
34
+
35
+ @classmethod
36
+ def generate(cls, bits: int = 1024, e: int = 65537) -> "RSAKey":
37
+ half = bits // 2
38
+ p = nt.rand_prime(half)
39
+ q = nt.rand_prime(bits - half)
40
+ while q == p:
41
+ q = nt.rand_prime(bits - half)
42
+ return cls._from_pq(p, q, e)
43
+
44
+ @classmethod
45
+ def _from_pq(cls, p: int, q: int, e: int) -> "RSAKey":
46
+ n = p * q
47
+ phi = (p - 1) * (q - 1)
48
+ d = nt.modinv(e, phi)
49
+ return cls(n=n, e=e, d=d, p=p, q=q)
50
+
51
+ @classmethod
52
+ def from_pq(cls, p: int, q: int, e: int = 65537) -> "RSAKey":
53
+ return cls._from_pq(p, q, e)
54
+
55
+ @classmethod
56
+ def from_ned(cls, n: int, e: int, d: int) -> "RSAKey":
57
+ return cls(n=n, e=e, d=d)
58
+
59
+ def public(self) -> "RSAKey":
60
+ return RSAKey(n=self.n, e=self.e)
61
+
62
+ def _size(self) -> int:
63
+ return (self.n.bit_length() + 7) // 8
64
+
65
+ def encrypt(self, message: bytes) -> bytes:
66
+ m = _b2i(message)
67
+ if m >= self.n:
68
+ raise ValueError("message too long for this RSA modulus")
69
+ return _i2b(pow(m, self.e, self.n), self._size())
70
+
71
+ def decrypt(self, ciphertext: bytes) -> bytes:
72
+ if self.d is None:
73
+ raise ValueError("cannot decrypt with a public-only key")
74
+ c = _b2i(ciphertext)
75
+ m = _i2b(pow(c, self.d, self.n), self._size()).lstrip(b"\x00")
76
+ return m or b"\x00"
77
+
78
+ def sign(self, message: bytes) -> bytes:
79
+ return self.decrypt(message)
80
+
81
+ def verify(self, message: bytes, signature: bytes) -> bool:
82
+ return _b2i(self.encrypt(signature)) == _b2i(message)
83
+
84
+ def max_bytes(self) -> int:
85
+ return (self.n.bit_length() - 1) // 8
86
+
87
+
88
+ def generate_keypair(bits: int = 1024, e: int = 65537) -> RSAKey:
89
+ return RSAKey.generate(bits, e)
90
+
91
+
92
+ def public_encrypt(n: int, e: int, message: bytes) -> bytes:
93
+ return RSAKey(n=n, e=e).encrypt(message)
94
+
95
+
96
+ def private_decrypt(n: int, e: int, d: int, ciphertext: bytes) -> bytes:
97
+ return RSAKey(n=n, e=e, d=d).decrypt(ciphertext)
98
+
99
+
100
+ def construct_private(n: int, p: int, q: int, e: int = 65537) -> RSAKey:
101
+ """Build the private key from a recovered factorisation."""
102
+ if p * q != n:
103
+ raise ValueError("p * q does not equal n")
104
+ return RSAKey._from_pq(p, q, e)
105
+
106
+
107
+ __all__ = ["RSAKey", "generate_keypair", "public_encrypt", "private_decrypt",
108
+ "construct_private"]
File without changes
@@ -0,0 +1,214 @@
1
+ """Deterministic AES attack recipes.
2
+
3
+ Two classic oracle attacks on misconfigured AES, both deterministic and
4
+ non-brute-force over the key space: they call an oracle once per secret
5
+ byte and reconstruct plaintext byte by byte.
6
+
7
+ - :func:`ecb_byte_at_a_time` — recovers the unknown suffix of an ECB
8
+ encryption when the attacker controls data placed between a fixed
9
+ (unknown) ``prefix`` and the ``secret``::
10
+
11
+ oracle(data) == ECB(prefix + data + secret)
12
+
13
+ Handles an arbitrary prefix length automatically.
14
+
15
+ - :func:`cbc_padding_oracle` — recovers plaintext block by block using a
16
+ padding-validity oracle on CBC ciphertexts::
17
+
18
+ oracle(previous || block) -> True if decrypting ``block`` against
19
+ ``previous`` yields valid PKCS#7 padding
20
+ """
21
+
22
+ from __future__ import annotations
23
+
24
+ BLOCK = 16
25
+
26
+
27
+ def _chunks(data: bytes, bs: int):
28
+ return [data[i * bs:(i + 1) * bs] for i in range(len(data) // bs)]
29
+
30
+
31
+ def _detect_blocksize(oracle) -> int:
32
+ base = len(oracle(b""))
33
+ prev = base
34
+ jumps = []
35
+ for n in range(1, 256):
36
+ cur = len(oracle(b"A" * n))
37
+ if cur != prev:
38
+ jumps.append(n)
39
+ prev = cur
40
+ if len(jumps) >= 2:
41
+ return jumps[1] - jumps[0]
42
+ raise ValueError("oracle output length never changes (no block size found)")
43
+
44
+
45
+ def _is_ecb(oracle, bs: int) -> bool:
46
+ out = oracle(b"A" * (bs * 3))
47
+ blocks = _chunks(out, bs)
48
+ return len(set(blocks)) != len(blocks)
49
+
50
+
51
+ def _find_prefix_align(oracle, bs: int):
52
+ """Return a list of candidate ``(q, j0)`` alignments.
53
+
54
+ Feeding ``A^q + X^(2*bs)`` (for two independent byte values X = B, C)
55
+ produces two consecutive identical ciphertext blocks iff
56
+ ``(prefix_len + q) % bs == 0``. The union of duplicate-pair indices
57
+ across the two characters makes the detector robust against a prefix
58
+ or secret that itself repeats bytes; the caller validates each
59
+ candidate by attempting a full recovery.
60
+ """
61
+ pair_sets = []
62
+ for byte_val in (0x42, 0x43):
63
+ pairs = set()
64
+ for q in range(bs):
65
+ out = oracle(b"A" * q + bytes([byte_val]) * (bs * 2))
66
+ blocks = _chunks(out, bs)
67
+ for i in range(len(blocks) - 1):
68
+ if blocks[i] == blocks[i + 1]:
69
+ pairs.add((q, i))
70
+ pair_sets.append(pairs)
71
+ cands = sorted(pair_sets[0] & pair_sets[1])
72
+ if not cands:
73
+ raise ValueError("could not align unknown prefix (no aligned run)")
74
+ return cands
75
+
76
+
77
+ def _plaintext_len(oracle, bs: int) -> int:
78
+ """Measure the oracle's total unpadded plaintext length.
79
+
80
+ Let M be the plaintext length for an empty attacker input. The oracle
81
+ output for x attacker bytes is ``bs * ceil((M + x) / bs)`` (a full
82
+ padding block when M + x is block-aligned, as PKCS#7 always pads), so
83
+ the smallest x at which the length grows reveals ``M % bs``.
84
+ """
85
+ w = len(oracle(b"")) // bs
86
+ prev = len(oracle(b""))
87
+ jump = None
88
+ for x in range(1, bs + 1):
89
+ cur = len(oracle(b"A" * x))
90
+ if cur > prev:
91
+ jump = x
92
+ break
93
+ prev = cur
94
+ if jump is None:
95
+ raise ValueError("oracle length never changes; cannot detect plaintext len")
96
+ return (w - 1) * bs + (bs - jump) % bs
97
+
98
+
99
+ def ecb_byte_at_a_time(oracle, secret_len=None, block_size=None):
100
+ """Recover ``secret`` from ``oracle(data) == ECB(prefix+data+secret)``.
101
+
102
+ Detects block size and ECB mode, measures an arbitrary unknown
103
+ ``prefix`` (alignment + length), and reconstructs ``secret`` byte by
104
+ byte with up to 256 oracle calls per byte plus a constant budget.
105
+ """
106
+ bs = _detect_blocksize(oracle) if block_size is None else block_size
107
+ if not _is_ecb(oracle, bs):
108
+ raise ValueError("oracle does not look like ECB (no duplicated block)")
109
+
110
+ total_len = _plaintext_len(oracle, bs)
111
+ candidates = _find_prefix_align(oracle, bs)
112
+ failures = []
113
+ best = None
114
+ for q, j0 in candidates:
115
+ try:
116
+ prefix_len = j0 * bs - q
117
+ if prefix_len < 0:
118
+ raise ValueError("negative prefix length")
119
+ slen = secret_len
120
+ if slen is None:
121
+ slen = total_len - prefix_len
122
+ if slen < 0:
123
+ raise ValueError("negative secret length")
124
+ recovered = _recover_with_alignment(oracle, bs, q, j0, slen)
125
+ except ValueError as exc:
126
+ failures.append((q, j0, str(exc)))
127
+ continue
128
+ if best is None or len(recovered) > len(best):
129
+ best = recovered
130
+ if best is None:
131
+ raise ValueError(
132
+ "could not recover secret with any alignment (%s)" % failures[0][2])
133
+ return best
134
+
135
+
136
+ def _recover_with_alignment(oracle, bs, q, j0, secret_len):
137
+ """Run the byte-at-a-time recovery assuming aligned offset (q, j0).
138
+
139
+ Reconstructs exactly ``secret_len`` bytes. A candidate alignment is
140
+ rejected when the probe byte never lands inside the compared block,
141
+ which would silently recover zero bytes instead of the secret.
142
+ """
143
+ known = bytearray()
144
+ for n in range(secret_len):
145
+ a, r = divmod(n, bs)
146
+ pad = bs - 1 - r
147
+ base = oracle(b"A" * (q + pad))
148
+ block = (j0 + a) * bs
149
+ if block + bs > len(base):
150
+ raise ValueError("target block out of range at byte %d" % n)
151
+ target_block = base[block:block + bs]
152
+ if n == 0:
153
+ # Probe two values first: the probe byte sits at a fixed
154
+ # position inside the compared block for a given alignment, so
155
+ # a match for both means the probe never reached that block at
156
+ # all (the alignment is degenerate) and the recovery would
157
+ # silently produce zero bytes.
158
+ zero = oracle(b"A" * (q + pad) + bytes([0]))
159
+ one = oracle(b"A" * (q + pad) + bytes([1]))
160
+ if (zero[block:block + bs] == target_block
161
+ and one[block:block + bs] == target_block):
162
+ raise ValueError(
163
+ "ambiguous alignment q=%d,j0=%d" % (q, j0))
164
+ found = None
165
+ for c in range(256):
166
+ cand = oracle(b"A" * (q + pad) + bytes(known) + bytes([c]))
167
+ if cand[block:block + bs] == target_block:
168
+ found = c
169
+ break
170
+ if found is None:
171
+ raise ValueError(
172
+ "byte %d not found with alignment q=%d,j0=%d"
173
+ % (len(known), q, j0))
174
+ known.append(found)
175
+ return bytes(known)
176
+
177
+
178
+ def cbc_padding_oracle(block_decrypt, iv, ciphertext, block_size=BLOCK):
179
+ """Recover the plaintext of a CBC ``ciphertext``.
180
+
181
+ ``block_decrypt`` is the padding oracle: given ``previous || block`` it
182
+ returns True iff decrypting ``block`` against ``previous`` yields a
183
+ block with valid PKCS#7 padding. No key material is required.
184
+
185
+ Runs the classic block-by-block attack: each block is recovered by
186
+ learning, byte by byte, the crafted previous block that makes the
187
+ padding valid, then XORing the result with the real IV (or previous
188
+ cipher block).
189
+ """
190
+ bs = block_size
191
+ if len(ciphertext) % bs:
192
+ raise ValueError("ciphertext must be a multiple of block size")
193
+ ct_blocks = _chunks(ciphertext, bs)
194
+ previous = iv
195
+ out = bytearray()
196
+ for block in ct_blocks:
197
+ recovered = bytearray(bs)
198
+ probe = bytearray(bs)
199
+ for pos in range(bs - 1, -1, -1):
200
+ pad_val = bs - pos
201
+ for i in range(pos + 1, bs):
202
+ probe[i] = recovered[i] ^ pad_val
203
+ found = None
204
+ for cand in range(256):
205
+ probe[pos] = cand
206
+ if block_decrypt(bytes(probe) + block):
207
+ found = cand
208
+ break
209
+ if found is None:
210
+ raise ValueError("padding oracle exhausted at pos %d" % pos)
211
+ recovered[pos] = found ^ pad_val
212
+ out += bytes(r ^ p for r, p in zip(recovered, previous))
213
+ previous = block
214
+ return bytes(out)
@@ -0,0 +1,134 @@
1
+ """Deterministic RSA attack recipes.
2
+
3
+ These recover factors or plaintext from a *misconfigured* RSA setup
4
+ without brute-forcing the key space. Every attack here either succeeds on
5
+ its precondition or returns ``None``; nothing here tries random exponents.
6
+
7
+ Implemented:
8
+ - :func:`wiener` — recovers ``d`` from a small private exponent.
9
+ - :func:`fermat` — factorial for close primes ``p``, ``q``.
10
+ - :func:`pollard_p_minus_1`— smooth ``p-1``.
11
+ - :func:`common_modulus` — two ciphertexts of the same message, public
12
+ exponents coprime.
13
+ - :func:`broadcast` — Håstad, low exponent + many receivers.
14
+ - :func:`plaintext_encrypt` — no-op composition helper.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ from crypto_gu import number_theory as nt
20
+
21
+
22
+ def _cf_frac(num: int, den: int) -> list:
23
+ """Continued fraction expansion of the rational num/den as ints."""
24
+ terms = []
25
+ while den:
26
+ terms.append(num // den)
27
+ num, den = den, num % den
28
+ return terms
29
+
30
+
31
+ def _convergents(terms: list) -> list:
32
+ out = []
33
+ n0, d0 = 0, 1
34
+ n1, d1 = 1, 0
35
+ for a in terms:
36
+ n0, n1 = n1, a * n1 + n0
37
+ d0, d1 = d1, a * d1 + d0
38
+ out.append((n1, d1))
39
+ return out
40
+
41
+
42
+ def wiener(n: int, e: int):
43
+ """Recover ``d`` (as an int) when it is smaller than ``n^0.25 / 3``.
44
+
45
+ Standard Wiener attack: the continued fraction of ``e/n`` contains
46
+ ``k/d``; test each convergent for a valid factorisation.
47
+ Returns the private exponent, or ``None``.
48
+ """
49
+ for k, d in _convergents(_cf_frac(e, n)):
50
+ if k == 0 or d == 0:
51
+ continue
52
+ if (e * d - 1) % k:
53
+ continue
54
+ phi = (e * d - 1) // k
55
+ s = n - phi + 1
56
+ disc = s * s - 4 * n
57
+ if disc < 0:
58
+ continue
59
+ r = nt.isqrt(disc)
60
+ if r * r != disc:
61
+ continue
62
+ p = (s + r) // 2
63
+ q = (s - r) // 2
64
+ if p * q == n:
65
+ return d
66
+ return None
67
+
68
+
69
+ def fermat(n: int, limit: int = 1000000):
70
+ """Recover factors when ``p`` and ``q`` are close (difference of squares)."""
71
+ return nt.fermat_factor(n, limit)
72
+
73
+
74
+ def pollard_p_minus_1(n: int, bound: int = 100000, rng=None):
75
+ """Recover a factor when ``p-1`` is bound-smooth."""
76
+ return nt.pollard_p_minus_1(n, bound, rng)
77
+
78
+
79
+ def common_modulus(n: int, e1: int, c1: int, e2: int, c2: int):
80
+ """Recover plaintext ``m`` (int) from two ciphertexts sharing modulus.
81
+
82
+ Precondition: ``gcd(e1, e2) == 1``. Returns the message int, or
83
+ ``None`` if the ciphertexts are inconsistent.
84
+ """
85
+ g, x, y = nt.extended_gcd(e1, e2)
86
+ if g != 1:
87
+ return None
88
+ if x < 0:
89
+ c1 = nt.modinv(c1, n)
90
+ x = -x
91
+ if y < 0:
92
+ c2 = nt.modinv(c2, n)
93
+ y = -y
94
+ return (pow(c1, x, n) * pow(c2, y, n)) % n
95
+
96
+
97
+ def broadcast(ciphertexts, exponent: int = 3):
98
+ """Håstad's broadcast attack.
99
+
100
+ ``ciphertexts`` is an iterable of ``(n, c)`` pairs for the same
101
+ plaintext encrypted to ``len(ciphertexts)`` >= ``exponent`` receivers.
102
+ Returns the message int, or ``None`` when precondition is unmet.
103
+ """
104
+ pairs = list(ciphertexts)
105
+ if len(pairs) < exponent:
106
+ return None
107
+ res = nt.crt([(c % n, n) for n, c in pairs])
108
+ return _iroot(res, exponent)
109
+
110
+
111
+ def _iroot(num: int, root: int) -> int:
112
+ """Integer ``root``-th root (floor) via Newton iterations."""
113
+ if root == 1:
114
+ return num
115
+ if num < 0:
116
+ return -_iroot(-num, root)
117
+ if num == 0:
118
+ return 0
119
+ high = 1 << ((num.bit_length() + root - 1) // root)
120
+ while True:
121
+ low = ((root - 1) * high + num // pow(high, root - 1)) // root
122
+ if low >= high:
123
+ break
124
+ high = low
125
+ return high
126
+
127
+
128
+ def plaintext_encrypt(key, message: bytes) -> bytes:
129
+ """Helper to compute a ciphertext from :class:`RSAKey` without padding."""
130
+ return key.encrypt(message)
131
+
132
+
133
+ __all__ = ["wiener", "fermat", "pollard_p_minus_1", "common_modulus",
134
+ "broadcast"]