lib/quic/src/crypto/packet.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const quic = @import("../root.zig");
  3 
  4 const crypto = quic.crypto;
  5 const Aes128Gcm = std.crypto.aead.aes_gcm.Aes128Gcm;
  6 const ChaCha20Poly1305 = std.crypto.aead.chacha_poly.ChaCha20Poly1305;
  7 
  8 pub const tag_bytes: usize = 16;
  9 /// The largest QUIC packet in bytes, which RFC 9000 section 18.2 sets as the largest UDP payload at
 10 /// 65,527, so a caller sizing a datagram buffer takes the largest packet this code will seal or
 11 /// open. Sealing or opening past it gives `PacketTooLarge`.
 12 pub const packet_bytes_max: usize = 65_527;
 13 
 14 const SealFailure = error{
 15     ConfidentialityLimitReached,
 16     InvalidHeader,
 17     InvalidLength,
 18     PacketCounterOverflow,
 19     PacketTooLarge,
 20     PacketTooShort,
 21 };
 22 
 23 const OpenFailure = error{
 24     AuthenticationFailed,
 25     IntegrityLimitReached,
 26     InvalidHeader,
 27     InvalidLength,
 28     InvalidPacketNumber,
 29     PacketTooLarge,
 30     PacketTooShort,
 31     ReservedBits,
 32     ScratchTooShort,
 33 };
 34 
 35 pub const SealError = crypto.header.ApplyError || SealFailure;
 36 pub const OpenError = quic.packet.DecodeError || crypto.header.ApplyError || OpenFailure;
 37 
 38 pub const Header = struct {
 39     first_byte: u8,
 40     pn_len: u3,
 41     pn_offset: usize,
 42     packet_end: usize,
 43     packet_number: u62,
 44     key_phase: ?bool,
 45 };
 46 
 47 pub const Opened = struct {
 48     packet_number: u62,
 49     payload: []u8,
 50     key_phase: ?bool,
 51 };
 52 
 53 /// Returns the 12-byte AEAD nonce for one packet number under one initialization vector, as RFC
 54 /// 9001 section 5.3 constructs it, so a caller sealing or opening a packet by hand computes the
 55 /// same nonce the peer will compute. The packet number is written big-endian into the low eight
 56 /// bytes of a zero field the width of the vector, and the two are exclusive-ored. Distinct packet
 57 /// numbers give distinct nonces under one vector, so a key stays usable across packets.
 58 pub fn nonce(iv: [crypto.iv_bytes]u8, packet_number: u62) [crypto.iv_bytes]u8 {
 59     var result = iv;
 60     var encoded: [8]u8 = undefined;
 61     std.mem.writeInt(u64, &encoded, packet_number, .big);
 62     for (0..8) |index| result[4 + index] ^= encoded[index];
 63     return result;
 64 }
 65 
 66 fn longEnd(packet: []const u8, pn_offset: usize) OpenError!usize {
 67     const decoded = try quic.packet.decodeLong(packet);
 68     const Metadata = struct { offset: usize, length: u62 };
 69     const metadata: Metadata = switch (decoded) {
 70         .initial => |value| .{ .offset = value.packet_number_offset, .length = value.length },
 71         .zero_rtt, .handshake => |value| .{
 72             .offset = value.packet_number_offset,
 73             .length = value.length,
 74         },
 75         .retry, .version_negotiation => return error.InvalidHeader,
 76     };
 77     if (metadata.offset != pn_offset) return error.InvalidHeader;
 78     const length = std.math.cast(usize, metadata.length) orelse
 79         return error.InvalidLength;
 80     if (pn_offset > packet.len) return error.InvalidLength;
 81     if (length > packet.len - pn_offset) return error.InvalidLength;
 82     return pn_offset + length;
 83 }
 84 
 85 fn protectedEnd(packet: []const u8, pn_offset: usize) OpenError!usize {
 86     if (packet.len == 0) return error.PacketTooShort;
 87     if (pn_offset == 0) return error.InvalidHeader;
 88     if (packet[0] & 0x80 != 0) {
 89         const packet_end = try longEnd(packet, pn_offset);
 90         if (packet_end > packet_bytes_max) return error.PacketTooLarge;
 91         return packet_end;
 92     }
 93     if (packet[0] & 0x40 == 0) return error.InvalidHeader;
 94     if (pn_offset > packet.len) return error.InvalidHeader;
 95     if (packet.len > packet_bytes_max) return error.PacketTooLarge;
 96     return packet.len;
 97 }
 98 
 99 fn sampleAt(packet: []const u8, pn_offset: usize, end: usize) OpenError![16]u8 {
100     const sample_offset = std.math.add(usize, pn_offset, 4) catch
101         return error.PacketTooShort;
102     const sample_end = std.math.add(usize, sample_offset, 16) catch
103         return error.PacketTooShort;
104     if (sample_end > end) return error.PacketTooShort;
105     return packet[sample_offset..][0..16].*;
106 }
107 
108 fn encrypt(
109     keys: *const crypto.Keys,
110     payload: []u8,
111     aad: []const u8,
112     packet_nonce: [12]u8,
113 ) [tag_bytes]u8 {
114     var tag: [tag_bytes]u8 = undefined;
115     const packet_key = keys.packetKey();
116     switch (keys.selectedSuite()) {
117         .aes_128_gcm_sha256 => Aes128Gcm.encrypt(
118             payload,
119             &tag,
120             payload,
121             aad,
122             packet_nonce,
123             packet_key[0..16].*,
124         ),
125         .chacha20_poly1305_sha256 => ChaCha20Poly1305.encrypt(
126             payload,
127             &tag,
128             payload,
129             aad,
130             packet_nonce,
131             packet_key,
132         ),
133     }
134     return tag;
135 }
136 
137 pub fn seal(
138     keys: *crypto.Keys,
139     packet_number: u62,
140     packet: []u8,
141     pn_offset: usize,
142     pn_len: u3,
143     payload_len: u16,
144 ) SealError!void {
145     if (keys.confidentialityExhausted()) return error.ConfidentialityLimitReached;
146     if (keys.sealed_packets == std.math.maxInt(u64)) return error.PacketCounterOverflow;
147     if (packet.len == 0) return error.PacketTooShort;
148     if (pn_len == 0 or pn_len > 4) return error.InvalidLength;
149     if (quic.packet.packetNumberLength(packet[0]) != pn_len) return error.InvalidHeader;
150     const header_len = std.math.add(usize, pn_offset, pn_len) catch
151         return error.InvalidLength;
152     const payload_bytes: usize = payload_len;
153     const payload_end = std.math.add(usize, header_len, payload_bytes) catch
154         return error.InvalidLength;
155     const packet_end = std.math.add(usize, payload_end, tag_bytes) catch
156         return error.InvalidLength;
157     if (packet_end > packet.len) return error.PacketTooShort;
158     if (packet_end > packet_bytes_max) return error.PacketTooLarge;
159     if (packet[0] & 0x80 != 0) {
160         const encoded_end = longEnd(packet, pn_offset) catch return error.InvalidHeader;
161         if (encoded_end != packet_end) return error.InvalidLength;
162     } else if (packet[0] & 0x40 == 0) return error.InvalidHeader;
163     std.debug.assert(packetNumberMatches(packet[pn_offset..header_len], packet_number));
164     const sample_offset = std.math.add(usize, pn_offset, 4) catch
165         return error.PacketTooShort;
166     const sample_end = std.math.add(usize, sample_offset, 16) catch
167         return error.PacketTooShort;
168     if (sample_end > packet_end) return error.PacketTooShort;
169     const payload = packet[header_len..payload_end];
170     const packet_nonce = nonce(keys.initializationVector(), packet_number);
171     const tag = encrypt(keys, payload, packet[0..header_len], packet_nonce);
172     packet[payload_end..][0..tag_bytes].* = tag;
173     std.debug.assert(sample_end <= packet_end);
174     const protected_sample = packet[sample_offset..][0..16].*;
175     const protection_mask = crypto.header.mask(keys, &protected_sample);
176     try crypto.header.protect(protection_mask, packet[0..packet_end], pn_offset, pn_len);
177     keys.sealed_packets += 1;
178 }
179 
180 fn packetNumberMatches(bytes: []const u8, packet_number: u62) bool {
181     std.debug.assert(bytes.len >= 1);
182     std.debug.assert(bytes.len <= 4);
183     var encoded: [8]u8 = undefined;
184     std.mem.writeInt(u64, &encoded, packet_number, .big);
185     return std.mem.eql(u8, bytes, encoded[encoded.len - bytes.len ..]);
186 }
187 
188 test "RFC 9001 section 5.3 seal packet number assertion contract" {
189     try std.testing.expect(packetNumberMatches(&.{0x34}, 0x1234));
190     try std.testing.expect(packetNumberMatches(&.{ 0x12, 0x34 }, 0x1234));
191     try std.testing.expect(!packetNumberMatches(&.{0x35}, 0x1234));
192 }
193 
194 fn readProtectedPacketNumber(
195     bytes: []const u8,
196     protection_mask: [crypto.header.mask_bytes]u8,
197 ) u32 {
198     std.debug.assert(bytes.len >= 1);
199     std.debug.assert(bytes.len <= 4);
200     var result: u32 = 0;
201     for (0..4) |index| {
202         if (index >= bytes.len) break;
203         result = (result << 8) | (bytes[index] ^ protection_mask[index + 1]);
204     }
205     return result;
206 }
207 
208 fn decrypt(
209     keys: *const crypto.Keys,
210     plaintext: []u8,
211     ciphertext: []const u8,
212     tag: [tag_bytes]u8,
213     aad: []const u8,
214     packet_nonce: [12]u8,
215 ) error{AuthenticationFailed}!void {
216     std.debug.assert(plaintext.len == ciphertext.len);
217     const packet_key = keys.packetKey();
218     switch (keys.selectedSuite()) {
219         .aes_128_gcm_sha256 => try Aes128Gcm.decrypt(
220             plaintext,
221             ciphertext,
222             tag,
223             aad,
224             packet_nonce,
225             packet_key[0..16].*,
226         ),
227         .chacha20_poly1305_sha256 => try ChaCha20Poly1305.decrypt(
228             plaintext,
229             ciphertext,
230             tag,
231             aad,
232             packet_nonce,
233             packet_key,
234         ),
235     }
236 }
237 
238 fn keyPhase(first_byte: u8) ?bool {
239     if (first_byte & 0x80 != 0) return null;
240     return first_byte & 0x04 != 0;
241 }
242 
243 pub fn unprotect(
244     keys: *const crypto.Keys,
245     packet: []u8,
246     pn_offset: usize,
247     largest_acked: ?u62,
248 ) OpenError!Header {
249     const packet_end = try protectedEnd(packet, pn_offset);
250     const sample = try sampleAt(packet, pn_offset, packet_end);
251     const protection_mask = crypto.header.mask(keys, &sample);
252     const inspected = try crypto.header.inspect(protection_mask, packet[0..packet_end], pn_offset);
253     const header_len = pn_offset + inspected.pn_len;
254     if (packet_end - header_len < tag_bytes) return error.InvalidLength;
255     if (largest_acked == std.math.maxInt(u62)) return error.InvalidPacketNumber;
256     const truncated = readProtectedPacketNumber(
257         packet[pn_offset..header_len],
258         protection_mask,
259     );
260     const bits: u6 = @as(u6, inspected.pn_len) * 8;
261     const packet_number = quic.packet.Number.expand(truncated, bits, largest_acked);
262     const pn_len = try crypto.header.unprotect(
263         protection_mask,
264         packet[0..packet_end],
265         pn_offset,
266     );
267     std.debug.assert(pn_len == inspected.pn_len);
268     return .{
269         .first_byte = inspected.first_byte,
270         .pn_len = pn_len,
271         .pn_offset = pn_offset,
272         .packet_end = packet_end,
273         .packet_number = packet_number,
274         .key_phase = keyPhase(inspected.first_byte),
275     };
276 }
277 
278 const PayloadLayout = struct {
279     header_len: usize,
280     tag_offset: usize,
281 };
282 
283 fn payloadLayout(packet: []const u8, parsed: Header) OpenError!PayloadLayout {
284     if (packet.len == 0) return error.PacketTooShort;
285     if (parsed.pn_len == 0 or parsed.pn_len > 4) return error.InvalidLength;
286     if (parsed.pn_offset == 0) return error.InvalidHeader;
287     if (parsed.pn_offset > parsed.packet_end) return error.InvalidHeader;
288     if (parsed.pn_len > parsed.packet_end - parsed.pn_offset) return error.InvalidLength;
289     const packet_end = try protectedEnd(packet, parsed.pn_offset);
290     if (packet_end != parsed.packet_end) return error.InvalidHeader;
291     const header_len = parsed.pn_offset + parsed.pn_len;
292     if (parsed.packet_end - header_len < tag_bytes) return error.InvalidLength;
293     if (packet[0] != parsed.first_byte) return error.InvalidHeader;
294     if (!packetNumberMatches(packet[parsed.pn_offset..header_len], parsed.packet_number)) {
295         return error.InvalidPacketNumber;
296     }
297     if (keyPhase(parsed.first_byte) != parsed.key_phase) return error.InvalidHeader;
298     return .{ .header_len = header_len, .tag_offset = parsed.packet_end - tag_bytes };
299 }
300 
301 fn slicesOverlap(first: []const u8, second: []const u8) bool {
302     if (first.len == 0 or second.len == 0) return false;
303     const first_address = @intFromPtr(first.ptr);
304     const second_address = @intFromPtr(second.ptr);
305     if (first_address <= second_address) {
306         return second_address - first_address < first.len;
307     }
308     return first_address - second_address < second.len;
309 }
310 
311 fn reservedBitsSet(first_byte: u8) bool {
312     const reserved_mask: u8 = if (first_byte & 0x80 != 0) 0x0c else 0x18;
313     return first_byte & reserved_mask != 0;
314 }
315 
316 fn openPayloadInner(
317     keys: *crypto.Keys,
318     packet: []u8,
319     parsed: Header,
320     scratch: []u8,
321 ) OpenError!Opened {
322     const layout = try payloadLayout(packet, parsed);
323     const ciphertext = packet[layout.header_len..layout.tag_offset];
324     if (scratch.len < ciphertext.len) return error.ScratchTooShort;
325     const plaintext = scratch[0..ciphertext.len];
326     std.debug.assert(!slicesOverlap(packet[0..parsed.packet_end], plaintext));
327     const tag: [tag_bytes]u8 = packet[layout.tag_offset..][0..tag_bytes].*;
328     try decrypt(
329         keys,
330         plaintext,
331         ciphertext,
332         tag,
333         packet[0..layout.header_len],
334         nonce(keys.initializationVector(), parsed.packet_number),
335     );
336     if (reservedBitsSet(parsed.first_byte)) return error.ReservedBits;
337     @memcpy(packet[layout.header_len..layout.tag_offset], plaintext);
338     return .{
339         .packet_number = parsed.packet_number,
340         .payload = packet[layout.header_len..layout.tag_offset],
341         .key_phase = parsed.key_phase,
342     };
343 }
344 
345 /// Decrypts the payload of a packet whose header `unprotect` has already parsed, and leaves the
346 /// plaintext in the packet where the ciphertext was, so a receiver decrypts the payload in place.
347 /// The caller supplies scratch bytes for the plaintext, and a scratch slice shorter than the
348 /// ciphertext gives `ScratchTooShort`. The ciphertext runs from the end of the packet number to the
349 /// authentication tag, so scratch needs the packet end less the packet number offset, the packet
350 /// number length, and the tag bytes. Scratch must lie outside the packet bytes, and only debug
351 /// builds check that. A failed authentication counts against the key's integrity limit, and a key
352 /// that has reached its limit gives `IntegrityLimitReached` before any decryption.
353 pub fn openPayload(
354     keys: *crypto.Keys,
355     packet: []u8,
356     parsed: Header,
357     scratch: []u8,
358 ) OpenError!Opened {
359     if (keys.exhausted()) return error.IntegrityLimitReached;
360     return openPayloadInner(keys, packet, parsed, scratch) catch |failure| {
361         if (failure != error.AuthenticationFailed) return failure;
362         std.debug.assert(keys.failed_opens < keys.integrityLimit());
363         keys.failed_opens += 1;
364         return failure;
365     };
366 }
367 
368 fn restoreHeader(keys: *const crypto.Keys, packet: []u8, parsed: Header) void {
369     const sample = sampleAt(packet, parsed.pn_offset, parsed.packet_end) catch unreachable;
370     const protection_mask = crypto.header.mask(keys, &sample);
371     crypto.header.protect(
372         protection_mask,
373         packet[0..parsed.packet_end],
374         parsed.pn_offset,
375         parsed.pn_len,
376     ) catch unreachable;
377 }
378 
379 pub fn open(
380     keys: *crypto.Keys,
381     packet: []u8,
382     scratch: []u8,
383     pn_offset: usize,
384     largest_acked: ?u62,
385 ) OpenError!Opened {
386     if (keys.exhausted()) return error.IntegrityLimitReached;
387     const parsed = try unprotect(keys, packet, pn_offset, largest_acked);
388     return openPayload(keys, packet, parsed, scratch) catch |failure| {
389         restoreHeader(keys, packet, parsed);
390         return failure;
391     };
392 }