Skip to content

Commit 2fdcde7

Browse files
dilaraacetinalxkm
andauthored
Feature/lwe encryption (#7631)
* feat(ciphers): add LWE (Regev) encryption * test(ciphers): add tests for LWE encryption --------- Co-authored-by: Oleksandr Klymenko <19151554+alxkm@users.noreply.github.com>
1 parent 922c4d3 commit 2fdcde7

2 files changed

Lines changed: 338 additions & 0 deletions

File tree

Lines changed: 216 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,216 @@
1+
package com.thealgorithms.ciphers;
2+
3+
import java.security.SecureRandom;
4+
5+
/**
6+
* Public-key encryption based on the Learning With Errors (LWE) problem, as proposed by Regev (2005).
7+
*
8+
* <p><b>Toy parameters for educational purposes only — NOT secure for real-world use.</b>
9+
*
10+
* <p>The LWE problem: given many noisy samples {@code b = A·s + e (mod q)} with a random matrix
11+
* {@code A} and small random errors {@code e}, it is hard to recover the secret vector {@code s}.
12+
* Without the errors, {@code s} could be found with Gaussian elimination.
13+
*
14+
* <p>The public key is {@code (A, b)} and the private key is {@code s}. To encrypt a bit, a random
15+
* subset {@code r} of the samples is summed: {@code u = rᵀ·A} and {@code v = r·b + bit·⌊q/2⌋}.
16+
* Decryption computes {@code v - u·s = r·e + bit·⌊q/2⌋}: a value close to 0 means bit 0, a value
17+
* close to {@code q/2} means bit 1.
18+
*
19+
* <p>Decryption is always correct because the remaining noise is bounded: {@code r} is a 0/1 vector
20+
* and {@code |e_i| ≤ B}, so {@code |r·e| ≤ m·B = 384 < q/4 ≈ 832}.
21+
*
22+
* <p>The modulus {@code q = 3329} is the one used by Kyber / ML-KEM (FIPS 203), which uses the
23+
* module-lattice version of this construction with discrete Gaussian-like noise and much larger
24+
* parameters. For simplicity, one instance holds both the public and the private key; in real use
25+
* they are kept separately.
26+
*
27+
* <p>Reference: <a href="https://en.wikipedia.org/wiki/Learning_with_errors">Wikipedia: Learning with errors</a>
28+
*
29+
* @author dilaraacetin
30+
*/
31+
public final class LWEEncryption {
32+
33+
// dimension of the secret vector
34+
private static final int N = 32;
35+
// modulus, same as Kyber / ML-KEM
36+
private static final int Q = 3329;
37+
// number of samples in the public key
38+
private static final int M = 128;
39+
// errors are uniform in [-B, B]
40+
private static final int B = 3;
41+
private static final int BITS_PER_BYTE = 8;
42+
43+
private final SecureRandom random = new SecureRandom();
44+
private final int[][] a;
45+
private final int[] b;
46+
private final int[] s;
47+
48+
/**
49+
* Generates a new key pair.
50+
*
51+
* @throws IllegalStateException if the parameters do not guarantee correct decryption
52+
*/
53+
public LWEEncryption() {
54+
checkDecryptionBound(M, B, Q);
55+
s = new int[N];
56+
for (int j = 0; j < N; j++) {
57+
s[j] = random.nextInt(Q);
58+
}
59+
a = new int[M][N];
60+
b = new int[M];
61+
for (int i = 0; i < M; i++) {
62+
long sum = 0;
63+
for (int j = 0; j < N; j++) {
64+
a[i][j] = random.nextInt(Q);
65+
sum += (long) a[i][j] * s[j];
66+
}
67+
int error = random.nextInt(2 * B + 1) - B;
68+
// b_i = A_i·s + e_i (mod q)
69+
b[i] = (int) Math.floorMod(sum + error, (long) Q);
70+
}
71+
}
72+
73+
/**
74+
* Encrypts a single bit. Encryption is randomized, so the same bit gives different ciphertexts.
75+
*
76+
* @param bit the bit to encrypt, 0 or 1
77+
* @return the ciphertext {@code (u, v)}
78+
* @throws IllegalArgumentException if {@code bit} is not 0 or 1
79+
*/
80+
public Ciphertext encryptBit(int bit) {
81+
if (bit != 0 && bit != 1) {
82+
throw new IllegalArgumentException("bit must be 0 or 1, got " + bit);
83+
}
84+
long[] u = new long[N];
85+
long v = (long) bit * (Q / 2);
86+
for (int i = 0; i < M; i++) {
87+
// r_i is 0 or 1: add sample i or skip it
88+
if (random.nextBoolean()) {
89+
for (int j = 0; j < N; j++) {
90+
u[j] += a[i][j];
91+
}
92+
v += b[i];
93+
}
94+
}
95+
int[] reduced = new int[N];
96+
for (int j = 0; j < N; j++) {
97+
reduced[j] = (int) Math.floorMod(u[j], (long) Q);
98+
}
99+
return new Ciphertext(reduced, (int) Math.floorMod(v, (long) Q));
100+
}
101+
102+
/**
103+
* Decrypts a single bit.
104+
*
105+
* @param ciphertext the ciphertext to decrypt
106+
* @return the decrypted bit, 0 or 1
107+
* @throws IllegalArgumentException if the ciphertext is null
108+
*/
109+
public int decryptBit(Ciphertext ciphertext) {
110+
if (ciphertext == null) {
111+
throw new IllegalArgumentException("ciphertext must not be null");
112+
}
113+
long dot = 0;
114+
for (int j = 0; j < N; j++) {
115+
dot += (long) ciphertext.u[j] * s[j];
116+
}
117+
// d = r·e + bit·⌊q/2⌋ (mod q); floorMod keeps it non-negative
118+
int d = (int) Math.floorMod(ciphertext.v - dot, (long) Q);
119+
return d >= Q / 4 && d < 3 * Q / 4 ? 1 : 0;
120+
}
121+
122+
/**
123+
* Encrypts a message bit by bit, most significant bit of each byte first.
124+
*
125+
* @param message the message to encrypt
126+
* @return eight ciphertexts per message byte
127+
* @throws IllegalArgumentException if the message is null or empty
128+
*/
129+
public Ciphertext[] encrypt(byte[] message) {
130+
if (message == null || message.length == 0) {
131+
throw new IllegalArgumentException("message must not be null or empty");
132+
}
133+
Ciphertext[] ciphertexts = new Ciphertext[message.length * BITS_PER_BYTE];
134+
for (int i = 0; i < message.length; i++) {
135+
for (int k = 0; k < BITS_PER_BYTE; k++) {
136+
ciphertexts[i * BITS_PER_BYTE + k] = encryptBit((message[i] >> (7 - k)) & 1);
137+
}
138+
}
139+
return ciphertexts;
140+
}
141+
142+
/**
143+
* Decrypts a message encrypted with {@link #encrypt(byte[])}.
144+
*
145+
* @param ciphertexts eight ciphertexts per message byte
146+
* @return the decrypted message
147+
* @throws IllegalArgumentException if the array is null, empty, contains null, or its length is
148+
* not a multiple of 8
149+
*/
150+
public byte[] decrypt(Ciphertext[] ciphertexts) {
151+
if (ciphertexts == null || ciphertexts.length == 0) {
152+
throw new IllegalArgumentException("ciphertexts must not be null or empty");
153+
}
154+
if (ciphertexts.length % BITS_PER_BYTE != 0) {
155+
throw new IllegalArgumentException("ciphertexts length must be a multiple of 8, got " + ciphertexts.length);
156+
}
157+
byte[] message = new byte[ciphertexts.length / BITS_PER_BYTE];
158+
for (int i = 0; i < message.length; i++) {
159+
int value = 0;
160+
for (int k = 0; k < BITS_PER_BYTE; k++) {
161+
value = (value << 1) | decryptBit(ciphertexts[i * BITS_PER_BYTE + k]);
162+
}
163+
message[i] = (byte) value;
164+
}
165+
return message;
166+
}
167+
168+
/**
169+
* Checks that the maximum decryption noise {@code m·B} stays below {@code q/4}.
170+
*
171+
* @throws IllegalStateException if decryption could fail with these parameters
172+
*/
173+
static void checkDecryptionBound(int samples, int errorBound, int modulus) {
174+
if (samples * errorBound >= modulus / 4) {
175+
throw new IllegalStateException("m * B must be less than q / 4 for correct decryption");
176+
}
177+
}
178+
179+
/**
180+
* An LWE ciphertext {@code (u, v)} for one bit. Immutable: the array is copied on construction
181+
* and on access.
182+
*/
183+
public static final class Ciphertext {
184+
private final int[] u;
185+
private final int v;
186+
187+
/**
188+
* Creates a ciphertext.
189+
*
190+
* @param u the vector part, of length 32
191+
* @param v the scalar part
192+
* @throws IllegalArgumentException if {@code u} is null or does not have length 32
193+
*/
194+
public Ciphertext(int[] u, int v) {
195+
if (u == null || u.length != N) {
196+
throw new IllegalArgumentException("u must contain exactly " + N + " values");
197+
}
198+
this.u = u.clone();
199+
this.v = v;
200+
}
201+
202+
/**
203+
* @return a copy of the vector part
204+
*/
205+
public int[] getU() {
206+
return u.clone();
207+
}
208+
209+
/**
210+
* @return the scalar part
211+
*/
212+
public int getV() {
213+
return v;
214+
}
215+
}
216+
}
Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,122 @@
1+
package com.thealgorithms.ciphers;
2+
3+
import static org.junit.jupiter.api.Assertions.assertArrayEquals;
4+
import static org.junit.jupiter.api.Assertions.assertEquals;
5+
import static org.junit.jupiter.api.Assertions.assertFalse;
6+
import static org.junit.jupiter.api.Assertions.assertThrows;
7+
8+
import java.nio.charset.StandardCharsets;
9+
import java.security.SecureRandom;
10+
import java.util.Arrays;
11+
import org.junit.jupiter.api.Test;
12+
import org.junit.jupiter.params.ParameterizedTest;
13+
import org.junit.jupiter.params.provider.ValueSource;
14+
15+
class LWEEncryptionTest {
16+
17+
@ParameterizedTest
18+
@ValueSource(ints = {0, 1})
19+
void testBitRoundTrip(int bit) {
20+
LWEEncryption lwe = new LWEEncryption();
21+
22+
for (int i = 0; i < 100; i++) {
23+
assertEquals(bit, lwe.decryptBit(lwe.encryptBit(bit)));
24+
}
25+
}
26+
27+
@Test
28+
void testTextMessageRoundTrip() {
29+
LWEEncryption lwe = new LWEEncryption();
30+
byte[] message = "Hello, PQC!".getBytes(StandardCharsets.UTF_8);
31+
32+
assertArrayEquals(message, lwe.decrypt(lwe.encrypt(message)));
33+
}
34+
35+
@Test
36+
void testRandomMessageRoundTrip() {
37+
LWEEncryption lwe = new LWEEncryption();
38+
byte[] message = new byte[1024];
39+
new SecureRandom().nextBytes(message);
40+
41+
assertArrayEquals(message, lwe.decrypt(lwe.encrypt(message)));
42+
}
43+
44+
@Test
45+
void testEncryptionIsRandomized() {
46+
LWEEncryption lwe = new LWEEncryption();
47+
48+
LWEEncryption.Ciphertext first = lwe.encryptBit(1);
49+
LWEEncryption.Ciphertext second = lwe.encryptBit(1);
50+
51+
assertFalse(Arrays.equals(toArray(first), toArray(second)));
52+
}
53+
54+
@Test
55+
void testWrongKeyDoesNotDecrypt() {
56+
LWEEncryption sender = new LWEEncryption();
57+
LWEEncryption other = new LWEEncryption();
58+
byte[] message = "sixteen byte msg".getBytes(StandardCharsets.UTF_8);
59+
60+
byte[] decrypted = other.decrypt(sender.encrypt(message));
61+
62+
assertFalse(Arrays.equals(message, decrypted));
63+
}
64+
65+
@ParameterizedTest
66+
@ValueSource(bytes = {(byte) 0b10110001, (byte) 0xFF, (byte) 0x80, 0x00, 0x01, 0x7F})
67+
void testSingleByteBitOrder(byte value) {
68+
LWEEncryption lwe = new LWEEncryption();
69+
70+
assertArrayEquals(new byte[] {value}, lwe.decrypt(lwe.encrypt(new byte[] {value})));
71+
}
72+
73+
@Test
74+
void testInvalidInput() {
75+
LWEEncryption lwe = new LWEEncryption();
76+
LWEEncryption.Ciphertext[] sevenCiphertexts = Arrays.copyOf(lwe.encrypt(new byte[] {1}), 7);
77+
LWEEncryption.Ciphertext[] withNull = lwe.encrypt(new byte[] {1});
78+
withNull[3] = null;
79+
80+
assertThrows(IllegalArgumentException.class, () -> lwe.encryptBit(2));
81+
assertThrows(IllegalArgumentException.class, () -> lwe.encryptBit(-1));
82+
assertThrows(IllegalArgumentException.class, () -> lwe.encrypt(null));
83+
assertThrows(IllegalArgumentException.class, () -> lwe.encrypt(new byte[0]));
84+
assertThrows(IllegalArgumentException.class, () -> lwe.decryptBit(null));
85+
assertThrows(IllegalArgumentException.class, () -> lwe.decrypt(null));
86+
assertThrows(IllegalArgumentException.class, () -> lwe.decrypt(new LWEEncryption.Ciphertext[0]));
87+
assertThrows(IllegalArgumentException.class, () -> lwe.decrypt(sevenCiphertexts));
88+
assertThrows(IllegalArgumentException.class, () -> lwe.decrypt(withNull));
89+
assertThrows(IllegalArgumentException.class, () -> new LWEEncryption.Ciphertext(null, 0));
90+
assertThrows(IllegalArgumentException.class, () -> new LWEEncryption.Ciphertext(new int[31], 0));
91+
assertThrows(IllegalArgumentException.class, () -> new LWEEncryption.Ciphertext(new int[33], 0));
92+
}
93+
94+
@Test
95+
void testCiphertextIsImmutable() {
96+
int[] u = new int[32];
97+
Arrays.fill(u, 7);
98+
LWEEncryption.Ciphertext ciphertext = new LWEEncryption.Ciphertext(u, 5);
99+
int[] expected = u.clone();
100+
101+
u[0] = 100;
102+
ciphertext.getU()[1] = 200;
103+
104+
assertArrayEquals(expected, ciphertext.getU());
105+
assertEquals(5, ciphertext.getV());
106+
}
107+
108+
@Test
109+
void testDecryptionBoundCheck() {
110+
LWEEncryption.checkDecryptionBound(128, 3, 3329);
111+
112+
assertThrows(IllegalStateException.class, () -> LWEEncryption.checkDecryptionBound(128, 7, 3329));
113+
}
114+
115+
// u followed by v, so that two ciphertexts can be compared with a single Arrays.equals
116+
private static int[] toArray(LWEEncryption.Ciphertext ciphertext) {
117+
int[] u = ciphertext.getU();
118+
int[] values = Arrays.copyOf(u, u.length + 1);
119+
values[u.length] = ciphertext.getV();
120+
return values;
121+
}
122+
}

0 commit comments

Comments
 (0)