authorgravatar for 124872+jedisct1@users.noreply.github.comFrank Denis <124872+jedisct1@users.noreply.github.com> 2023-05-23 09:55:45+02:00
committergravatar for noreply@github.comGitHub <noreply@github.com> 2023-05-23 09:55:45+02:00
log9d179a98f69dbab393cbb3fc5dd4b64c553a721b
tree87ed4994d77e60cc82c3cc241710a861abd14edc
parenta0652fb93077322c331345e1aeff25d5338f30b0
signaturebadge-question-mark Signed by PGP key 4AEE18F83AFDEB23

Make Poly1305 faster by leveraging @addWithOverflow/@subWithOverflow (#15815)

These operations are constant-time on most, if not all currently supported architectures. However, even if they are not, this is not a big deal in the case on Poly1305, as the key is added at the end. The final addition remains protected. SalsaPoly and ChaChaPoly do encrypt-then-mac, so side channels would not leak anything about the plaintext anyway. * Apple Silicon (M1) Before: 2048 MiB/s After : 2823 MiB/s * AMD Ryzen 7 Before: 3165 MiB/s After : 4774 MiB/s

1 files changed, 90 insertions(+), 96 deletions(-)

lib/std/crypto/poly1305.zig+90-96
...@@ -1,6 +1,7 @@...@@ -1,6 +1,7 @@
1const std = @import("../std.zig");1const std = @import("../std.zig");
2const utils = std.crypto.utils;2const utils = std.crypto.utils;
3const mem = std.mem;3const mem = std.mem;
4const mulWide = std.math.mulWide;
45
5pub const Poly1305 = struct {6pub const Poly1305 = struct {
6 pub const block_length: usize = 16;7 pub const block_length: usize = 16;
...@@ -8,7 +9,7 @@ pub const Poly1305 = struct {...@@ -8,7 +9,7 @@ pub const Poly1305 = struct {
8 pub const key_length = 32;9 pub const key_length = 32;
910
10 // constant multiplier (from the secret key)11 // constant multiplier (from the secret key)
11 r: [3]u64,12 r: [2]u64,
12 // accumulated hash13 // accumulated hash
13 h: [3]u64 = [_]u64{ 0, 0, 0 },14 h: [3]u64 = [_]u64{ 0, 0, 0 },
14 // random number added at the end (from the secret key)15 // random number added at the end (from the secret key)
...@@ -19,13 +20,10 @@ pub const Poly1305 = struct {...@@ -19,13 +20,10 @@ pub const Poly1305 = struct {
19 buf: [block_length]u8 align(16) = undefined,20 buf: [block_length]u8 align(16) = undefined,
2021
21 pub fn init(key: *const [key_length]u8) Poly1305 {22 pub fn init(key: *const [key_length]u8) Poly1305 {
22 const t0 = mem.readIntLittle(u64, key[0..8]);
23 const t1 = mem.readIntLittle(u64, key[8..16]);
24 return Poly1305{23 return Poly1305{
25 .r = [_]u64{24 .r = [_]u64{
26 t0 & 0xffc0fffffff,25 mem.readIntLittle(u64, key[0..8]) & 0x0ffffffc0fffffff,
27 ((t0 >> 44) | (t1 << 20)) & 0xfffffc0ffff,26 mem.readIntLittle(u64, key[8..16]) & 0x0ffffffc0ffffffc,
28 ((t1 >> 24)) & 0x00ffffffc0f,
29 },27 },
30 .pad = [_]u64{28 .pad = [_]u64{
31 mem.readIntLittle(u64, key[16..24]),29 mem.readIntLittle(u64, key[16..24]),
...@@ -34,43 +32,77 @@ pub const Poly1305 = struct {...@@ -34,43 +32,77 @@ pub const Poly1305 = struct {
34 };32 };
35 }33 }
3634
35 inline fn add(a: u64, b: u64, c: u1) struct { u64, u1 } {
36 const v1 = @addWithOverflow(a, b);
37 const v2 = @addWithOverflow(v1[0], c);
38 return .{ v2[0], v1[1] | v2[1] };
39 }
40
41 inline fn sub(a: u64, b: u64, c: u1) struct { u64, u1 } {
42 const v1 = @subWithOverflow(a, b);
43 const v2 = @subWithOverflow(v1[0], c);
44 return .{ v2[0], v1[1] | v2[1] };
45 }
46
37 fn blocks(st: *Poly1305, m: []const u8, comptime last: bool) void {47 fn blocks(st: *Poly1305, m: []const u8, comptime last: bool) void {
38 const hibit: u64 = if (last) 0 else 1 << 40;48 const hibit: u64 = if (last) 0 else 1;
39 const r0 = st.r[0];49 const r0 = st.r[0];
40 const r1 = st.r[1];50 const r1 = st.r[1];
41 const r2 = st.r[2];51
42 var h0 = st.h[0];52 var h0 = st.h[0];
43 var h1 = st.h[1];53 var h1 = st.h[1];
44 var h2 = st.h[2];54 var h2 = st.h[2];
45 const s1 = r1 * (5 << 2);55
46 const s2 = r2 * (5 << 2);
47 var i: usize = 0;56 var i: usize = 0;
57
48 while (i + block_length <= m.len) : (i += block_length) {58 while (i + block_length <= m.len) : (i += block_length) {
49 // h += m[i]59 const in0 = mem.readIntLittle(u64, m[i..][0..8]);
50 const t0 = mem.readIntLittle(u64, m[i..][0..8]);60 const in1 = mem.readIntLittle(u64, m[i + 8 ..][0..8]);
51 const t1 = mem.readIntLittle(u64, m[i + 8 ..][0..8]);61
52 h0 += @truncate(u44, t0);62 // Add the input message to H
53 h1 += @truncate(u44, (t0 >> 44) | (t1 << 20));63 var v = @addWithOverflow(h0, in0);
54 h2 += @truncate(u42, t1 >> 24) | hibit;64 h0 = v[0];
5565 v = add(h1, in1, v[1]);
56 // h *= r66 h1 = v[0];
57 const d0 = @as(u128, h0) * r0 + @as(u128, h1) * s2 + @as(u128, h2) * s1;67 h2 +%= v[1] +% hibit;
58 var d1 = @as(u128, h0) * r1 + @as(u128, h1) * r0 + @as(u128, h2) * s2;68
59 var d2 = @as(u128, h0) * r2 + @as(u128, h1) * r1 + @as(u128, h2) * r0;69 // Compute H * R
6070 const m0 = mulWide(u64, h0, r0);
61 // partial reduction71 const h1r0 = mulWide(u64, h1, r0);
62 var carry = @intCast(u64, d0 >> 44);72 const h0r1 = mulWide(u64, h0, r1);
63 h0 = @truncate(u44, d0);73 const h2r0 = mulWide(u64, h2, r0);
64 d1 += carry;74 const h1r1 = mulWide(u64, h1, r1);
65 carry = @intCast(u64, d1 >> 44);75 const m3 = mulWide(u64, h2, r1);
66 h1 = @truncate(u44, d1);76 const m1 = h1r0 +% h0r1;
67 d2 += carry;77 const m2 = h2r0 +% h1r1;
68 carry = @intCast(u64, d2 >> 42);78
69 h2 = @truncate(u42, d2);79 const t0 = @truncate(u64, m0);
70 h0 += @truncate(u64, carry) * 5;80 v = @addWithOverflow(@truncate(u64, m1), @truncate(u64, m0 >> 64));
71 carry = h0 >> 44;81 const t1 = v[0];
72 h0 = @truncate(u44, h0);82 v = add(@truncate(u64, m2), @truncate(u64, m1 >> 64), v[1]);
73 h1 += carry;83 const t2 = v[0];
84 v = add(@truncate(u64, m3), @truncate(u64, m2 >> 64), v[1]);
85 const t3 = v[0];
86
87 // Partial reduction
88 h0 = t0;
89 h1 = t1;
90 h2 = t2 & 3;
91
92 // Add c*(4+1)
93 var cclo = t2 & ~@as(u64, 3);
94 var cchi = t3;
95 v = @addWithOverflow(h0, cclo);
96 h0 = v[0];
97 v = add(h1, cchi, v[1]);
98 h1 = v[0];
99 h2 +%= v[1];
100 const cc = (cclo | (@as(u128, cchi) << 64)) >> 2;
101 v = @addWithOverflow(h0, @truncate(u64, cc));
102 h0 = v[0];
103 v = add(h1, @truncate(u64, cc >> 64), v[1]);
104 h1 = v[0];
105 h2 +%= v[1];
74 }106 }
75 st.h = [_]u64{ h0, h1, h2 };107 st.h = [_]u64{ h0, h1, h2 };
76 }108 }
...@@ -115,10 +147,7 @@ pub const Poly1305 = struct {...@@ -115,10 +147,7 @@ pub const Poly1305 = struct {
115 if (st.leftover == 0) {147 if (st.leftover == 0) {
116 return;148 return;
117 }149 }
118 var i = st.leftover;150 @memset(st.buf[st.leftover..], 0);
119 while (i < block_length) : (i += 1) {
120 st.buf[i] = 0;
121 }
122 st.blocks(&st.buf);151 st.blocks(&st.buf);
123 st.leftover = 0;152 st.leftover = 0;
124 }153 }
...@@ -128,65 +157,30 @@ pub const Poly1305 = struct {...@@ -128,65 +157,30 @@ pub const Poly1305 = struct {
128 var i = st.leftover;157 var i = st.leftover;
129 st.buf[i] = 1;158 st.buf[i] = 1;
130 i += 1;159 i += 1;
131 while (i < block_length) : (i += 1) {160 @memset(st.buf[i..], 0);
132 st.buf[i] = 0;
133 }
134 st.blocks(&st.buf, true);161 st.blocks(&st.buf, true);
135 }162 }
136 // fully carry h163
137 var carry = st.h[1] >> 44;164 var h0 = st.h[0];
138 st.h[1] = @truncate(u44, st.h[1]);165 var h1 = st.h[1];
139 st.h[2] += carry;166 var h2 = st.h[2];
140 carry = st.h[2] >> 42;167
141 st.h[2] = @truncate(u42, st.h[2]);168 // H - (2^130 - 5)
142 st.h[0] += carry * 5;169 var v = sub(h0, 0xfffffffffffffffb, 0);
143 carry = st.h[0] >> 44;170 const h_p0 = v[0];
144 st.h[0] = @truncate(u44, st.h[0]);171 v = sub(h1, 0xffffffffffffffff, v[1]);
145 st.h[1] += carry;172 const h_p1 = v[0];
146 carry = st.h[1] >> 44;173 v = sub(h2, 0x0000000000000003, v[1]);
147 st.h[1] = @truncate(u44, st.h[1]);174
148 st.h[2] += carry;175 // Final reduction, subtract 2^130-5 from H if H >= 2^130-5
149 carry = st.h[2] >> 42;176 const mask = v[1] -% 1;
150 st.h[2] = @truncate(u42, st.h[2]);177 h0 ^= mask & (h0 ^ h_p0);
151 st.h[0] += carry * 5;178 h1 ^= mask & (h1 ^ h_p1);
152 carry = st.h[0] >> 44;179
153 st.h[0] = @truncate(u44, st.h[0]);180 // Add the first half of the key, we intentionally don't use @addWithOverflow() here.
154 st.h[1] += carry;181 st.h[0] = h0 +% st.pad[0];
155182 const c = ((h0 & st.pad[0]) | ((h0 | st.pad[0]) & ~st.h[0])) >> 63;
156 // compute h + -p183 st.h[1] = h1 +% st.pad[1] +% c;
157 var g0 = st.h[0] + 5;
158 carry = g0 >> 44;
159 g0 = @truncate(u44, g0);
160 var g1 = st.h[1] + carry;
161 carry = g1 >> 44;
162 g1 = @truncate(u44, g1);
163 var g2 = st.h[2] + carry -% (1 << 42);
164
165 // (hopefully) constant-time select h if h < p, or h + -p if h >= p
166 const mask = (g2 >> 63) -% 1;
167 g0 &= mask;
168 g1 &= mask;
169 g2 &= mask;
170 const nmask = ~mask;
171 st.h[0] = (st.h[0] & nmask) | g0;
172 st.h[1] = (st.h[1] & nmask) | g1;
173 st.h[2] = (st.h[2] & nmask) | g2;
174
175 // h = (h + pad)
176 const t0 = st.pad[0];
177 const t1 = st.pad[1];
178 st.h[0] += @truncate(u44, t0);
179 carry = st.h[0] >> 44;
180 st.h[0] = @truncate(u44, st.h[0]);
181 st.h[1] += @truncate(u44, (t0 >> 44) | (t1 << 20)) + carry;
182 carry = st.h[1] >> 44;
183 st.h[1] = @truncate(u44, st.h[1]);
184 st.h[2] += @truncate(u42, t1 >> 24) + carry;
185 st.h[2] = @truncate(u42, st.h[2]);
186
187 // mac = h % (2^128)
188 st.h[0] |= st.h[1] << 44;
189 st.h[1] = (st.h[1] >> 20) | (st.h[2] << 24);
190184
191 mem.writeIntLittle(u64, out[0..8], st.h[0]);185 mem.writeIntLittle(u64, out[0..8], st.h[0]);
192 mem.writeIntLittle(u64, out[8..16], st.h[1]);186 mem.writeIntLittle(u64, out[8..16], st.h[1]);