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 }