| ... | @@ -1,6 +1,7 @@ | ... | @@ -1,6 +1,7 @@ |
| 1 | const std = @import("../std.zig"); | 1 | const std = @import("../std.zig"); |
| 2 | const utils = std.crypto.utils; | 2 | const utils = std.crypto.utils; |
| 3 | const mem = std.mem; | 3 | const mem = std.mem; |
| | 4 | const mulWide = std.math.mulWide; |
| 4 | | 5 | |
| 5 | pub const Poly1305 = struct { | 6 | pub 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; |
| 9 | | 10 | |
| 10 | // constant multiplier (from the secret key) | 11 | // constant multiplier (from the secret key) |
| 11 | r: [3]u64, | 12 | r: [2]u64, |
| 12 | // accumulated hash | 13 | // 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, |
| 20 | | 21 | |
| 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 | } |
| 36 | | 34 | |
| | 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]; |
| 55 | | 65 | v = add(h1, in1, v[1]); |
| 56 | // h *= r | 66 | 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 |
| 60 | | 70 | const m0 = mulWide(u64, h0, r0); |
| 61 | // partial reduction | 71 | 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 h | 163 | |
| 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]; |
| 155 | | 182 | const c = ((h0 & st.pad[0]) | ((h0 | st.pad[0]) & ~st.h[0])) >> 63; |
| 156 | // compute h + -p | 183 | 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); | | |
| 190 | | 184 | |
| 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]); |