authorgravatar for 124872+jedisct1@users.noreply.github.comFrank Denis <124872+jedisct1@users.noreply.github.com> 2023-03-13 22:18:26+01:00
committergravatar for noreply@github.comGitHub <noreply@github.com> 2023-03-13 22:18:26+01:00
log962299157840979ba659d478785f5ed0759d5401
treebf1346750ed2a57b40442fcdda0fd54918f9dc90
parentd525ecb523fed3c1496bf2f5315d06ea95716c8b
signaturebadge-question-mark Signed by PGP key 4AEE18F83AFDEB23

Add configurable side channels mitigations; enable them on soft AES (#13739)

* Add configurable side channels mitigations; enable them on soft AES Our software AES implementation doesn't have any mitigations against side channels. Go's generic implementation is not protected at all either, and even OpenSSL only has minimal mitigations. Full mitigations against cache-based attacks (bitslicing, fixslicing) come at a huge performance cost, making AES-based primitives pretty much useless for many applications. They also don't offer any protection against other classes of side channel attacks. In practice, partially protected, or even unprotected implementations are not as bad as it sounds. Exploiting these side channels requires an attacker that is able to submit many plaintexts/ciphertexts and perform accurate measurements. Noisy measurements can still be exploited, but require a significant amount of attempts. Wether this is exploitable or not depends on the platform, application and the attacker's proximity. So, some libraries made the choice of minimal mitigations and some use better mitigations in spite of the performance hit. It's a tradeoff (security vs performance), and there's no one-size-fits all implementation. What applies to AES applies to other cryptographic primitives. For example, RSA signatures are very sensible to fault attacks, regardless of them using the CRT or not. A mitigation is to verify every produced signature. That also comes with a performance cost. Wether to do it or not depends on wether fault attacks are part of the threat model or not. Thanks to Zig's comptime, we can try to address these different requirements. This PR adds a `side_channels_protection` global, that can later be complemented with `fault_attacks_protection` and possibly other knobs. It can have 4 different values: - `none`: which doesn't enable additional mitigations. "Additional", because it only disables mitigations that don't have a big performance cost. For example, checking authentication tags will still be done in constant time. - `basic`: which enables mitigations protecting against attacks in a common scenario, where an attacker doesn't have physical access to the device, cannot run arbitrary code on the same thread, and cannot conduct brute-force attacks without being throttled. - `medium`: which enables additional mitigations, offering practical protection in a shared environement. - `full`: which enables all the mitigations we have. The tradeoff is that the more mitigations we enable, the bigger the performance hit will be. But this let applications choose what's best for their use case. `medium` is the default. Currently, this only affects software AES, but that setting can later be used by other primitives. For AES, our implementation is a traditional table-based, with 4 32-bit tables and a sbox. Lookups in that table have been replaced by function calls. These functions can add a configurable noise level, making cache-based attacks more difficult to conduct. In the `none` mitigation level, the behavior is exactly the same as before. Performance also remains the same. In other levels, we compress the T tables into a single one, and read data from multiple cache lines (all of them in `full` mode), for all bytes in parallel. More precise measurements and way more attempts become necessary in order to find correlations. In addition, we use distinct copies of the sbox for key expansion and encryption, so that they don't share the same L1 cache entries. The best known attacks target the first two AES round, or the last one. While future attacks may improve on this, AES achieves full diffusion after 4 rounds. So, we can relax the mitigations after that. This is what this implementation does, enabling mitigations again for the last two rounds. In `full` mode, all the rounds are protected. The protection assumes that lookups within a cache line are secret. The cachebleed attack showed that it can be circumvented, but that requires an attacker to be able to abuse hyperthreading and run code on the same core as the encryption, which is rarely a practical scenario. Still, the current AES API allows us to transparently switch to using fixslicing/bitslicing later when the `full` mitigation level is enabled. * Software AES: use little-endian representation. Virtually all platforms are little-endian these days, so optimizing for big-endian CPUs doesn't make sense any more.

2 files changed, 329 insertions(+), 63 deletions(-)

lib/std/crypto.zig+27
......@@ -1,3 +1,5 @@
1const root = @import("root");
2
13/// Authenticated Encryption with Associated Data
24pub const aead = struct {
35 pub const aegis = struct {
......@@ -183,6 +185,31 @@ pub const errors = @import("crypto/errors.zig");
183185pub const tls = @import("crypto/tls.zig");
184186pub const Certificate = @import("crypto/Certificate.zig");
185187
188/// Global configuration of cryptographic implementations in the standard library.
189pub const config = struct {
190 /// Side-channels mitigations.
191 pub const SideChannelsMitigations = enum {
192 /// No additional side-channel mitigations are applied.
193 /// This is the fastest mode.
194 none,
195 /// The `basic` mode protects against most practical attacks, provided that the
196 /// application or implements proper defenses against brute-force attacks.
197 /// It offers a good balance between performance and security.
198 basic,
199 /// The `medium` mode offers increased resilience against side-channel attacks,
200 /// making most attacks unpractical even on shared/low latency environements.
201 /// This is the default mode.
202 medium,
203 /// The `full` mode offers the highest level of protection against side-channel attacks.
204 /// Note that this doesn't cover all possible attacks (especially power analysis or
205 /// thread-local attacks such as cachebleed), and that the performance impact is significant.
206 full,
207 };
208
209 /// This is a global configuration that applies to all cryptographic implementations.
210 pub const side_channels_mitigations: SideChannelsMitigations = if (@hasDecl(root, "side_channels_mitigations")) root.side_channels_mitigations else .medium;
211};
212
186213test {
187214 _ = aead.aegis.Aegis128L;
188215 _ = aead.aegis.Aegis256;
lib/std/crypto/aes/soft.zig+302-63
......@@ -1,11 +1,11 @@
1// Based on Go stdlib implementation
2
31const std = @import("../../std.zig");
42const math = std.math;
53const mem = std.mem;
64
75const BlockVec = [4]u32;
86
7const side_channels_mitigations = std.crypto.config.side_channels_mitigations;
8
99/// A single AES block.
1010pub const Block = struct {
1111 pub const block_length: usize = 16;
......@@ -15,20 +15,20 @@ pub const Block = struct {
1515
1616 /// Convert a byte sequence into an internal representation.
1717 pub inline fn fromBytes(bytes: *const [16]u8) Block {
18 const s0 = mem.readIntBig(u32, bytes[0..4]);
19 const s1 = mem.readIntBig(u32, bytes[4..8]);
20 const s2 = mem.readIntBig(u32, bytes[8..12]);
21 const s3 = mem.readIntBig(u32, bytes[12..16]);
18 const s0 = mem.readIntLittle(u32, bytes[0..4]);
19 const s1 = mem.readIntLittle(u32, bytes[4..8]);
20 const s2 = mem.readIntLittle(u32, bytes[8..12]);
21 const s3 = mem.readIntLittle(u32, bytes[12..16]);
2222 return Block{ .repr = BlockVec{ s0, s1, s2, s3 } };
2323 }
2424
2525 /// Convert the internal representation of a block into a byte sequence.
2626 pub inline fn toBytes(block: Block) [16]u8 {
2727 var bytes: [16]u8 = undefined;
28 mem.writeIntBig(u32, bytes[0..4], block.repr[0]);
29 mem.writeIntBig(u32, bytes[4..8], block.repr[1]);
30 mem.writeIntBig(u32, bytes[8..12], block.repr[2]);
31 mem.writeIntBig(u32, bytes[12..16], block.repr[3]);
28 mem.writeIntLittle(u32, bytes[0..4], block.repr[0]);
29 mem.writeIntLittle(u32, bytes[4..8], block.repr[1]);
30 mem.writeIntLittle(u32, bytes[8..12], block.repr[2]);
31 mem.writeIntLittle(u32, bytes[12..16], block.repr[3]);
3232 return bytes;
3333 }
3434
......@@ -50,32 +50,93 @@ pub const Block = struct {
5050 const s2 = block.repr[2];
5151 const s3 = block.repr[3];
5252
53 const t0 = round_key.repr[0] ^ table_encrypt[0][@truncate(u8, s0 >> 24)] ^ table_encrypt[1][@truncate(u8, s1 >> 16)] ^ table_encrypt[2][@truncate(u8, s2 >> 8)] ^ table_encrypt[3][@truncate(u8, s3)];
54 const t1 = round_key.repr[1] ^ table_encrypt[0][@truncate(u8, s1 >> 24)] ^ table_encrypt[1][@truncate(u8, s2 >> 16)] ^ table_encrypt[2][@truncate(u8, s3 >> 8)] ^ table_encrypt[3][@truncate(u8, s0)];
55 const t2 = round_key.repr[2] ^ table_encrypt[0][@truncate(u8, s2 >> 24)] ^ table_encrypt[1][@truncate(u8, s3 >> 16)] ^ table_encrypt[2][@truncate(u8, s0 >> 8)] ^ table_encrypt[3][@truncate(u8, s1)];
56 const t3 = round_key.repr[3] ^ table_encrypt[0][@truncate(u8, s3 >> 24)] ^ table_encrypt[1][@truncate(u8, s0 >> 16)] ^ table_encrypt[2][@truncate(u8, s1 >> 8)] ^ table_encrypt[3][@truncate(u8, s2)];
53 var x: [4]u32 = undefined;
54 x = table_lookup(&table_encrypt, @truncate(u8, s0), @truncate(u8, s1 >> 8), @truncate(u8, s2 >> 16), @truncate(u8, s3 >> 24));
55 var t0 = x[0] ^ x[1] ^ x[2] ^ x[3];
56 x = table_lookup(&table_encrypt, @truncate(u8, s1), @truncate(u8, s2 >> 8), @truncate(u8, s3 >> 16), @truncate(u8, s0 >> 24));
57 var t1 = x[0] ^ x[1] ^ x[2] ^ x[3];
58 x = table_lookup(&table_encrypt, @truncate(u8, s2), @truncate(u8, s3 >> 8), @truncate(u8, s0 >> 16), @truncate(u8, s1 >> 24));
59 var t2 = x[0] ^ x[1] ^ x[2] ^ x[3];
60 x = table_lookup(&table_encrypt, @truncate(u8, s3), @truncate(u8, s0 >> 8), @truncate(u8, s1 >> 16), @truncate(u8, s2 >> 24));
61 var t3 = x[0] ^ x[1] ^ x[2] ^ x[3];
62
63 t0 ^= round_key.repr[0];
64 t1 ^= round_key.repr[1];
65 t2 ^= round_key.repr[2];
66 t3 ^= round_key.repr[3];
67
68 return Block{ .repr = BlockVec{ t0, t1, t2, t3 } };
69 }
70
71 /// Encrypt a block with a round key *WITHOUT ANY PROTECTION AGAINST SIDE CHANNELS*
72 pub inline fn encryptUnprotected(block: Block, round_key: Block) Block {
73 const s0 = block.repr[0];
74 const s1 = block.repr[1];
75 const s2 = block.repr[2];
76 const s3 = block.repr[3];
77
78 var x: [4]u32 = undefined;
79 x = .{
80 table_encrypt[0][@truncate(u8, s0)],
81 table_encrypt[1][@truncate(u8, s1 >> 8)],
82 table_encrypt[2][@truncate(u8, s2 >> 16)],
83 table_encrypt[3][@truncate(u8, s3 >> 24)],
84 };
85 var t0 = x[0] ^ x[1] ^ x[2] ^ x[3];
86 x = .{
87 table_encrypt[0][@truncate(u8, s1)],
88 table_encrypt[1][@truncate(u8, s2 >> 8)],
89 table_encrypt[2][@truncate(u8, s3 >> 16)],
90 table_encrypt[3][@truncate(u8, s0 >> 24)],
91 };
92 var t1 = x[0] ^ x[1] ^ x[2] ^ x[3];
93 x = .{
94 table_encrypt[0][@truncate(u8, s2)],
95 table_encrypt[1][@truncate(u8, s3 >> 8)],
96 table_encrypt[2][@truncate(u8, s0 >> 16)],
97 table_encrypt[3][@truncate(u8, s1 >> 24)],
98 };
99 var t2 = x[0] ^ x[1] ^ x[2] ^ x[3];
100 x = .{
101 table_encrypt[0][@truncate(u8, s3)],
102 table_encrypt[1][@truncate(u8, s0 >> 8)],
103 table_encrypt[2][@truncate(u8, s1 >> 16)],
104 table_encrypt[3][@truncate(u8, s2 >> 24)],
105 };
106 var t3 = x[0] ^ x[1] ^ x[2] ^ x[3];
107
108 t0 ^= round_key.repr[0];
109 t1 ^= round_key.repr[1];
110 t2 ^= round_key.repr[2];
111 t3 ^= round_key.repr[3];
57112
58113 return Block{ .repr = BlockVec{ t0, t1, t2, t3 } };
59114 }
60115
61116 /// Encrypt a block with the last round key.
62117 pub inline fn encryptLast(block: Block, round_key: Block) Block {
63 const t0 = block.repr[0];
64 const t1 = block.repr[1];
65 const t2 = block.repr[2];
66 const t3 = block.repr[3];
118 const s0 = block.repr[0];
119 const s1 = block.repr[1];
120 const s2 = block.repr[2];
121 const s3 = block.repr[3];
67122
68123 // Last round uses s-box directly and XORs to produce output.
69 var s0 = @as(u32, sbox_encrypt[t0 >> 24]) << 24 | @as(u32, sbox_encrypt[t1 >> 16 & 0xff]) << 16 | @as(u32, sbox_encrypt[t2 >> 8 & 0xff]) << 8 | @as(u32, sbox_encrypt[t3 & 0xff]);
70 var s1 = @as(u32, sbox_encrypt[t1 >> 24]) << 24 | @as(u32, sbox_encrypt[t2 >> 16 & 0xff]) << 16 | @as(u32, sbox_encrypt[t3 >> 8 & 0xff]) << 8 | @as(u32, sbox_encrypt[t0 & 0xff]);
71 var s2 = @as(u32, sbox_encrypt[t2 >> 24]) << 24 | @as(u32, sbox_encrypt[t3 >> 16 & 0xff]) << 16 | @as(u32, sbox_encrypt[t0 >> 8 & 0xff]) << 8 | @as(u32, sbox_encrypt[t1 & 0xff]);
72 var s3 = @as(u32, sbox_encrypt[t3 >> 24]) << 24 | @as(u32, sbox_encrypt[t0 >> 16 & 0xff]) << 16 | @as(u32, sbox_encrypt[t1 >> 8 & 0xff]) << 8 | @as(u32, sbox_encrypt[t2 & 0xff]);
73 s0 ^= round_key.repr[0];
74 s1 ^= round_key.repr[1];
75 s2 ^= round_key.repr[2];
76 s3 ^= round_key.repr[3];
124 var x: [4]u8 = undefined;
125 x = sbox_lookup(&sbox_encrypt, @truncate(u8, s3 >> 24), @truncate(u8, s2 >> 16), @truncate(u8, s1 >> 8), @truncate(u8, s0));
126 var t0 = @as(u32, x[0]) << 24 | @as(u32, x[1]) << 16 | @as(u32, x[2]) << 8 | @as(u32, x[3]);
127 x = sbox_lookup(&sbox_encrypt, @truncate(u8, s0 >> 24), @truncate(u8, s3 >> 16), @truncate(u8, s2 >> 8), @truncate(u8, s1));
128 var t1 = @as(u32, x[0]) << 24 | @as(u32, x[1]) << 16 | @as(u32, x[2]) << 8 | @as(u32, x[3]);
129 x = sbox_lookup(&sbox_encrypt, @truncate(u8, s1 >> 24), @truncate(u8, s0 >> 16), @truncate(u8, s3 >> 8), @truncate(u8, s2));
130 var t2 = @as(u32, x[0]) << 24 | @as(u32, x[1]) << 16 | @as(u32, x[2]) << 8 | @as(u32, x[3]);
131 x = sbox_lookup(&sbox_encrypt, @truncate(u8, s2 >> 24), @truncate(u8, s1 >> 16), @truncate(u8, s0 >> 8), @truncate(u8, s3));
132 var t3 = @as(u32, x[0]) << 24 | @as(u32, x[1]) << 16 | @as(u32, x[2]) << 8 | @as(u32, x[3]);
133
134 t0 ^= round_key.repr[0];
135 t1 ^= round_key.repr[1];
136 t2 ^= round_key.repr[2];
137 t3 ^= round_key.repr[3];
77138
78 return Block{ .repr = BlockVec{ s0, s1, s2, s3 } };
139 return Block{ .repr = BlockVec{ t0, t1, t2, t3 } };
79140 }
80141
81142 /// Decrypt a block with a round key.
......@@ -85,32 +146,93 @@ pub const Block = struct {
85146 const s2 = block.repr[2];
86147 const s3 = block.repr[3];
87148
88 const t0 = round_key.repr[0] ^ table_decrypt[0][@truncate(u8, s0 >> 24)] ^ table_decrypt[1][@truncate(u8, s3 >> 16)] ^ table_decrypt[2][@truncate(u8, s2 >> 8)] ^ table_decrypt[3][@truncate(u8, s1)];
89 const t1 = round_key.repr[1] ^ table_decrypt[0][@truncate(u8, s1 >> 24)] ^ table_decrypt[1][@truncate(u8, s0 >> 16)] ^ table_decrypt[2][@truncate(u8, s3 >> 8)] ^ table_decrypt[3][@truncate(u8, s2)];
90 const t2 = round_key.repr[2] ^ table_decrypt[0][@truncate(u8, s2 >> 24)] ^ table_decrypt[1][@truncate(u8, s1 >> 16)] ^ table_decrypt[2][@truncate(u8, s0 >> 8)] ^ table_decrypt[3][@truncate(u8, s3)];
91 const t3 = round_key.repr[3] ^ table_decrypt[0][@truncate(u8, s3 >> 24)] ^ table_decrypt[1][@truncate(u8, s2 >> 16)] ^ table_decrypt[2][@truncate(u8, s1 >> 8)] ^ table_decrypt[3][@truncate(u8, s0)];
149 var x: [4]u32 = undefined;
150 x = table_lookup(&table_decrypt, @truncate(u8, s0), @truncate(u8, s3 >> 8), @truncate(u8, s2 >> 16), @truncate(u8, s1 >> 24));
151 var t0 = x[0] ^ x[1] ^ x[2] ^ x[3];
152 x = table_lookup(&table_decrypt, @truncate(u8, s1), @truncate(u8, s0 >> 8), @truncate(u8, s3 >> 16), @truncate(u8, s2 >> 24));
153 var t1 = x[0] ^ x[1] ^ x[2] ^ x[3];
154 x = table_lookup(&table_decrypt, @truncate(u8, s2), @truncate(u8, s1 >> 8), @truncate(u8, s0 >> 16), @truncate(u8, s3 >> 24));
155 var t2 = x[0] ^ x[1] ^ x[2] ^ x[3];
156 x = table_lookup(&table_decrypt, @truncate(u8, s3), @truncate(u8, s2 >> 8), @truncate(u8, s1 >> 16), @truncate(u8, s0 >> 24));
157 var t3 = x[0] ^ x[1] ^ x[2] ^ x[3];
158
159 t0 ^= round_key.repr[0];
160 t1 ^= round_key.repr[1];
161 t2 ^= round_key.repr[2];
162 t3 ^= round_key.repr[3];
163
164 return Block{ .repr = BlockVec{ t0, t1, t2, t3 } };
165 }
166
167 /// Decrypt a block with a round key *WITHOUT ANY PROTECTION AGAINST SIDE CHANNELS*
168 pub inline fn decryptUnprotected(block: Block, round_key: Block) Block {
169 const s0 = block.repr[0];
170 const s1 = block.repr[1];
171 const s2 = block.repr[2];
172 const s3 = block.repr[3];
173
174 var x: [4]u32 = undefined;
175 x = .{
176 table_decrypt[0][@truncate(u8, s0)],
177 table_decrypt[1][@truncate(u8, s3 >> 8)],
178 table_decrypt[2][@truncate(u8, s2 >> 16)],
179 table_decrypt[3][@truncate(u8, s1 >> 24)],
180 };
181 var t0 = x[0] ^ x[1] ^ x[2] ^ x[3];
182 x = .{
183 table_decrypt[0][@truncate(u8, s1)],
184 table_decrypt[1][@truncate(u8, s0 >> 8)],
185 table_decrypt[2][@truncate(u8, s3 >> 16)],
186 table_decrypt[3][@truncate(u8, s2 >> 24)],
187 };
188 var t1 = x[0] ^ x[1] ^ x[2] ^ x[3];
189 x = .{
190 table_decrypt[0][@truncate(u8, s2)],
191 table_decrypt[1][@truncate(u8, s1 >> 8)],
192 table_decrypt[2][@truncate(u8, s0 >> 16)],
193 table_decrypt[3][@truncate(u8, s3 >> 24)],
194 };
195 var t2 = x[0] ^ x[1] ^ x[2] ^ x[3];
196 x = .{
197 table_decrypt[0][@truncate(u8, s3)],
198 table_decrypt[1][@truncate(u8, s2 >> 8)],
199 table_decrypt[2][@truncate(u8, s1 >> 16)],
200 table_decrypt[3][@truncate(u8, s0 >> 24)],
201 };
202 var t3 = x[0] ^ x[1] ^ x[2] ^ x[3];
203
204 t0 ^= round_key.repr[0];
205 t1 ^= round_key.repr[1];
206 t2 ^= round_key.repr[2];
207 t3 ^= round_key.repr[3];
92208
93209 return Block{ .repr = BlockVec{ t0, t1, t2, t3 } };
94210 }
95211
96212 /// Decrypt a block with the last round key.
97213 pub inline fn decryptLast(block: Block, round_key: Block) Block {
98 const t0 = block.repr[0];
99 const t1 = block.repr[1];
100 const t2 = block.repr[2];
101 const t3 = block.repr[3];
214 const s0 = block.repr[0];
215 const s1 = block.repr[1];
216 const s2 = block.repr[2];
217 const s3 = block.repr[3];
102218
103219 // Last round uses s-box directly and XORs to produce output.
104 var s0 = @as(u32, sbox_decrypt[t0 >> 24]) << 24 | @as(u32, sbox_decrypt[t3 >> 16 & 0xff]) << 16 | @as(u32, sbox_decrypt[t2 >> 8 & 0xff]) << 8 | @as(u32, sbox_decrypt[t1 & 0xff]);
105 var s1 = @as(u32, sbox_decrypt[t1 >> 24]) << 24 | @as(u32, sbox_decrypt[t0 >> 16 & 0xff]) << 16 | @as(u32, sbox_decrypt[t3 >> 8 & 0xff]) << 8 | @as(u32, sbox_decrypt[t2 & 0xff]);
106 var s2 = @as(u32, sbox_decrypt[t2 >> 24]) << 24 | @as(u32, sbox_decrypt[t1 >> 16 & 0xff]) << 16 | @as(u32, sbox_decrypt[t0 >> 8 & 0xff]) << 8 | @as(u32, sbox_decrypt[t3 & 0xff]);
107 var s3 = @as(u32, sbox_decrypt[t3 >> 24]) << 24 | @as(u32, sbox_decrypt[t2 >> 16 & 0xff]) << 16 | @as(u32, sbox_decrypt[t1 >> 8 & 0xff]) << 8 | @as(u32, sbox_decrypt[t0 & 0xff]);
108 s0 ^= round_key.repr[0];
109 s1 ^= round_key.repr[1];
110 s2 ^= round_key.repr[2];
111 s3 ^= round_key.repr[3];
220 var x: [4]u8 = undefined;
221 x = sbox_lookup(&sbox_decrypt, @truncate(u8, s1 >> 24), @truncate(u8, s2 >> 16), @truncate(u8, s3 >> 8), @truncate(u8, s0));
222 var t0 = @as(u32, x[0]) << 24 | @as(u32, x[1]) << 16 | @as(u32, x[2]) << 8 | @as(u32, x[3]);
223 x = sbox_lookup(&sbox_decrypt, @truncate(u8, s2 >> 24), @truncate(u8, s3 >> 16), @truncate(u8, s0 >> 8), @truncate(u8, s1));
224 var t1 = @as(u32, x[0]) << 24 | @as(u32, x[1]) << 16 | @as(u32, x[2]) << 8 | @as(u32, x[3]);
225 x = sbox_lookup(&sbox_decrypt, @truncate(u8, s3 >> 24), @truncate(u8, s0 >> 16), @truncate(u8, s1 >> 8), @truncate(u8, s2));
226 var t2 = @as(u32, x[0]) << 24 | @as(u32, x[1]) << 16 | @as(u32, x[2]) << 8 | @as(u32, x[3]);
227 x = sbox_lookup(&sbox_decrypt, @truncate(u8, s0 >> 24), @truncate(u8, s1 >> 16), @truncate(u8, s2 >> 8), @truncate(u8, s3));
228 var t3 = @as(u32, x[0]) << 24 | @as(u32, x[1]) << 16 | @as(u32, x[2]) << 8 | @as(u32, x[3]);
229
230 t0 ^= round_key.repr[0];
231 t1 ^= round_key.repr[1];
232 t2 ^= round_key.repr[2];
233 t3 ^= round_key.repr[3];
112234
113 return Block{ .repr = BlockVec{ s0, s1, s2, s3 } };
235 return Block{ .repr = BlockVec{ t0, t1, t2, t3 } };
114236 }
115237
116238 /// Apply the bitwise XOR operation to the content of two blocks.
......@@ -226,7 +348,8 @@ fn KeySchedule(comptime Aes: type) type {
226348 const subw = struct {
227349 // Apply sbox_encrypt to each byte in w.
228350 fn func(w: u32) u32 {
229 return @as(u32, sbox_encrypt[w >> 24]) << 24 | @as(u32, sbox_encrypt[w >> 16 & 0xff]) << 16 | @as(u32, sbox_encrypt[w >> 8 & 0xff]) << 8 | @as(u32, sbox_encrypt[w & 0xff]);
351 const x = sbox_lookup(&sbox_key_schedule, @truncate(u8, w), @truncate(u8, w >> 8), @truncate(u8, w >> 16), @truncate(u8, w >> 24));
352 return @as(u32, x[3]) << 24 | @as(u32, x[2]) << 16 | @as(u32, x[1]) << 8 | @as(u32, x[0]);
230353 }
231354 }.func;
232355
......@@ -244,6 +367,10 @@ fn KeySchedule(comptime Aes: type) type {
244367 }
245368 round_keys[i / 4].repr[i % 4] = round_keys[(i - words_in_key) / 4].repr[(i - words_in_key) % 4] ^ t;
246369 }
370 i = 0;
371 inline while (i < round_keys.len * 4) : (i += 1) {
372 round_keys[i / 4].repr[i % 4] = @byteSwap(round_keys[i / 4].repr[i % 4]);
373 }
247374 return Self{ .round_keys = round_keys };
248375 }
249376
......@@ -257,11 +384,13 @@ fn KeySchedule(comptime Aes: type) type {
257384 const ei = total_words - i - 4;
258385 comptime var j: usize = 0;
259386 inline while (j < 4) : (j += 1) {
260 var x = round_keys[(ei + j) / 4].repr[(ei + j) % 4];
387 var rk = round_keys[(ei + j) / 4].repr[(ei + j) % 4];
261388 if (i > 0 and i + 4 < total_words) {
262 x = table_decrypt[0][sbox_encrypt[x >> 24]] ^ table_decrypt[1][sbox_encrypt[x >> 16 & 0xff]] ^ table_decrypt[2][sbox_encrypt[x >> 8 & 0xff]] ^ table_decrypt[3][sbox_encrypt[x & 0xff]];
389 const x = sbox_lookup(&sbox_key_schedule, @truncate(u8, rk >> 24), @truncate(u8, rk >> 16), @truncate(u8, rk >> 8), @truncate(u8, rk));
390 const y = table_lookup(&table_decrypt, x[3], x[2], x[1], x[0]);
391 rk = y[0] ^ y[1] ^ y[2] ^ y[3];
263392 }
264 inv_round_keys[(i + j) / 4].repr[(i + j) % 4] = x;
393 inv_round_keys[(i + j) / 4].repr[(i + j) % 4] = rk;
265394 }
266395 }
267396 return Self{ .round_keys = inv_round_keys };
......@@ -293,7 +422,17 @@ pub fn AesEncryptCtx(comptime Aes: type) type {
293422 const round_keys = ctx.key_schedule.round_keys;
294423 var t = Block.fromBytes(src).xorBlocks(round_keys[0]);
295424 comptime var i = 1;
296 inline while (i < rounds) : (i += 1) {
425 if (side_channels_mitigations == .full) {
426 inline while (i < rounds) : (i += 1) {
427 t = t.encrypt(round_keys[i]);
428 }
429 } else {
430 inline while (i < 5) : (i += 1) {
431 t = t.encrypt(round_keys[i]);
432 }
433 inline while (i < rounds - 1) : (i += 1) {
434 t = t.encryptUnprotected(round_keys[i]);
435 }
297436 t = t.encrypt(round_keys[i]);
298437 }
299438 t = t.encryptLast(round_keys[rounds]);
......@@ -305,7 +444,17 @@ pub fn AesEncryptCtx(comptime Aes: type) type {
305444 const round_keys = ctx.key_schedule.round_keys;
306445 var t = Block.fromBytes(&counter).xorBlocks(round_keys[0]);
307446 comptime var i = 1;
308 inline while (i < rounds) : (i += 1) {
447 if (side_channels_mitigations == .full) {
448 inline while (i < rounds) : (i += 1) {
449 t = t.encrypt(round_keys[i]);
450 }
451 } else {
452 inline while (i < 5) : (i += 1) {
453 t = t.encrypt(round_keys[i]);
454 }
455 inline while (i < rounds - 1) : (i += 1) {
456 t = t.encryptUnprotected(round_keys[i]);
457 }
309458 t = t.encrypt(round_keys[i]);
310459 }
311460 t = t.encryptLast(round_keys[rounds]);
......@@ -359,7 +508,17 @@ pub fn AesDecryptCtx(comptime Aes: type) type {
359508 const inv_round_keys = ctx.key_schedule.round_keys;
360509 var t = Block.fromBytes(src).xorBlocks(inv_round_keys[0]);
361510 comptime var i = 1;
362 inline while (i < rounds) : (i += 1) {
511 if (side_channels_mitigations == .full) {
512 inline while (i < rounds) : (i += 1) {
513 t = t.decrypt(inv_round_keys[i]);
514 }
515 } else {
516 inline while (i < 5) : (i += 1) {
517 t = t.decrypt(inv_round_keys[i]);
518 }
519 inline while (i < rounds - 1) : (i += 1) {
520 t = t.decryptUnprotected(inv_round_keys[i]);
521 }
363522 t = t.decrypt(inv_round_keys[i]);
364523 }
365524 t = t.decryptLast(inv_round_keys[rounds]);
......@@ -428,10 +587,11 @@ const powx = init: {
428587 break :init array;
429588};
430589
431const sbox_encrypt align(64) = generateSbox(false);
432const sbox_decrypt align(64) = generateSbox(true);
433const table_encrypt align(64) = generateTable(false);
434const table_decrypt align(64) = generateTable(true);
590const sbox_encrypt align(64) = generateSbox(false); // S-box for encryption
591const sbox_key_schedule align(64) = generateSbox(false); // S-box only for key schedule, so that it uses distinct L1 cache entries than the S-box used for encryption
592const sbox_decrypt align(64) = generateSbox(true); // S-box for decryption
593const table_encrypt align(64) = generateTable(false); // 4-byte LUTs for encryption
594const table_decrypt align(64) = generateTable(true); // 4-byte LUTs for decryption
435595
436596// Generate S-box substitution values.
437597fn generateSbox(invert: bool) [256]u8 {
......@@ -472,14 +632,14 @@ fn generateTable(invert: bool) [4][256]u32 {
472632 var table: [4][256]u32 = undefined;
473633
474634 for (generateSbox(invert), 0..) |value, index| {
475 table[0][index] = mul(value, if (invert) 0xb else 0x3);
476 table[0][index] |= math.shl(u32, mul(value, if (invert) 0xd else 0x1), 8);
477 table[0][index] |= math.shl(u32, mul(value, if (invert) 0x9 else 0x1), 16);
478 table[0][index] |= math.shl(u32, mul(value, if (invert) 0xe else 0x2), 24);
479
480 table[1][index] = math.rotr(u32, table[0][index], 8);
481 table[2][index] = math.rotr(u32, table[0][index], 16);
482 table[3][index] = math.rotr(u32, table[0][index], 24);
635 table[0][index] = math.shl(u32, mul(value, if (invert) 0xb else 0x3), 24);
636 table[0][index] |= math.shl(u32, mul(value, if (invert) 0xd else 0x1), 16);
637 table[0][index] |= math.shl(u32, mul(value, if (invert) 0x9 else 0x1), 8);
638 table[0][index] |= mul(value, if (invert) 0xe else 0x2);
639
640 table[1][index] = math.rotl(u32, table[0][index], 8);
641 table[2][index] = math.rotl(u32, table[0][index], 16);
642 table[3][index] = math.rotl(u32, table[0][index], 24);
483643 }
484644
485645 return table;
......@@ -506,3 +666,82 @@ fn mul(a: u8, b: u8) u8 {
506666
507667 return @truncate(u8, s);
508668}
669
670const cache_line_bytes = 64;
671
672inline fn sbox_lookup(sbox: *align(64) const [256]u8, idx0: u8, idx1: u8, idx2: u8, idx3: u8) [4]u8 {
673 if (side_channels_mitigations == .none) {
674 return [4]u8{
675 sbox[idx0],
676 sbox[idx1],
677 sbox[idx2],
678 sbox[idx3],
679 };
680 } else {
681 const stride = switch (side_channels_mitigations) {
682 .none => unreachable,
683 .basic => sbox.len / 4,
684 .medium => sbox.len / (sbox.len / cache_line_bytes) * 2,
685 .full => sbox.len / (sbox.len / cache_line_bytes),
686 };
687 const of0 = idx0 % stride;
688 const of1 = idx1 % stride;
689 const of2 = idx2 % stride;
690 const of3 = idx3 % stride;
691 var t: [4][sbox.len / stride]u8 align(64) = undefined;
692 var i: usize = 0;
693 while (i < t[0].len) : (i += 1) {
694 const tx = sbox[i * stride ..];
695 t[0][i] = tx[of0];
696 t[1][i] = tx[of1];
697 t[2][i] = tx[of2];
698 t[3][i] = tx[of3];
699 }
700 std.mem.doNotOptimizeAway(t);
701 return [4]u8{
702 t[0][idx0 / stride],
703 t[1][idx1 / stride],
704 t[2][idx2 / stride],
705 t[3][idx3 / stride],
706 };
707 }
708}
709
710inline fn table_lookup(table: *align(64) const [4][256]u32, idx0: u8, idx1: u8, idx2: u8, idx3: u8) [4]u32 {
711 if (side_channels_mitigations == .none) {
712 return [4]u32{
713 table[0][idx0],
714 table[1][idx1],
715 table[2][idx2],
716 table[3][idx3],
717 };
718 } else {
719 const table_bytes = @sizeOf(@TypeOf(table[0]));
720 const stride = switch (side_channels_mitigations) {
721 .none => unreachable,
722 .basic => table[0].len / 4,
723 .medium => table[0].len / (table_bytes / cache_line_bytes) * 2,
724 .full => table[0].len / (table_bytes / cache_line_bytes),
725 };
726 const of0 = idx0 % stride;
727 const of1 = idx1 % stride;
728 const of2 = idx2 % stride;
729 const of3 = idx3 % stride;
730 var t: [4][table[0].len / stride]u32 align(64) = undefined;
731 var i: usize = 0;
732 while (i < t[0].len) : (i += 1) {
733 const tx = table[0][i * stride ..];
734 t[0][i] = tx[of0];
735 t[1][i] = tx[of1];
736 t[2][i] = tx[of2];
737 t[3][i] = tx[of3];
738 }
739 std.mem.doNotOptimizeAway(t);
740 return [4]u32{
741 t[0][idx0 / stride],
742 math.rotl(u32, t[1][idx1 / stride], 8),
743 math.rotl(u32, t[2][idx2 / stride], 16),
744 math.rotl(u32, t[3][idx3 / stride], 24),
745 };
746 }
747}