lib/quic/src/tls/engine/extension.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const quic = @import("../../root.zig");
3 const tls = @import("../root.zig");
4
5 const cursor = quic.cursor;
6 const ExtensionType = std.crypto.tls.ExtensionType;
7
8 const ExtensionSpecific = error{
9 IllegalExtension,
10 InvalidExtension,
11 UnsupportedExtension,
12 };
13 pub const Error = tls.message.DecodeError || ExtensionSpecific;
14
15 /// Bounds every list scan in this file so a peer cannot make one of them run long. The largest
16 /// number of entries this parser reads from one extension list is 256, and every list scan here
17 /// runs at most that many times, so the work one extension can cost is fixed in advance. The bound
18 /// belongs to the narrow subset of TLS 1.3 this engine encodes and accepts, fixed message by
19 /// message, its TLS profile. This profile is narrower than what the wire format allows, so the RFC
20 /// does not state this limit. A key share, protocol name, or server name list that runs past the
21 /// bound gives `InvalidExtension`.
22 const profile_list_entries_max: usize = std.math.maxInt(u8) + 1;
23
24 pub const ClientOffer = struct {
25 tls_1_3: bool = false,
26 x25519_group: bool = false,
27 ed25519_signature: bool = false,
28 key_share: ?[32]u8 = null,
29 alpn: ?[]const u8 = null,
30 transport_parameters: ?[]const u8 = null,
31 server_raw_key: bool = false,
32 client_raw_key: bool = false,
33 server_name: ?[]const u8 = null,
34 };
35
36 pub const ServerSelection = struct {
37 tls_1_3: bool = false,
38 key_share: ?[32]u8 = null,
39 };
40
41 pub const EncryptedSelection = struct {
42 alpn: ?[]const u8 = null,
43 transport_parameters: ?[]const u8 = null,
44 server_raw_key: bool = false,
45 client_raw_key: bool = false,
46 };
47
48 pub fn clientOffer(extensions: tls.message.Extensions, wanted_alpn: []const u8) Error!ClientOffer {
49 var result = ClientOffer{};
50 var iterator = extensions.iterator();
51 var seen: [tls.message.extension_count_max]u16 = undefined;
52 var count: u8 = 0;
53 for (0..tls.message.extension_count_max + 1) |_| {
54 const extension = try iterator.next() orelse return result;
55 try mark(&seen, &count, extension.kind);
56 switch (extension.kind) {
57 .supported_versions => result.tls_1_3 = try hasU16(extension.data, 0x0304, true),
58 .supported_groups => result.x25519_group = try hasU16(
59 extension.data,
60 @backingInt(std.crypto.tls.NamedGroup.x25519),
61 false,
62 ),
63 .signature_algorithms => result.ed25519_signature = try hasU16(
64 extension.data,
65 @backingInt(std.crypto.tls.SignatureScheme.ed25519),
66 false,
67 ),
68 .key_share => result.key_share = try clientKeyShare(extension.data),
69 .application_layer_protocol_negotiation => {
70 result.alpn = try selectAlpn(extension.data, wanted_alpn);
71 },
72 .quic_transport_parameters => result.transport_parameters = extension.data,
73 .server_certificate_type => result.server_raw_key = try offersRawKey(extension.data),
74 .client_certificate_type => result.client_raw_key = try offersRawKey(extension.data),
75 .server_name => result.server_name = try serverName(extension.data),
76 else => if (known(extension.kind) and !clientHelloAllowed(extension.kind)) {
77 return error.IllegalExtension;
78 },
79 }
80 }
81 unreachable;
82 }
83
84 pub fn serverSelection(extensions: tls.message.Extensions) Error!ServerSelection {
85 var result = ServerSelection{};
86 var iterator = extensions.iterator();
87 var seen: [tls.message.extension_count_max]u16 = undefined;
88 var count: u8 = 0;
89 for (0..tls.message.extension_count_max + 1) |_| {
90 const extension = try iterator.next() orelse return result;
91 try mark(&seen, &count, extension.kind);
92 switch (extension.kind) {
93 .supported_versions => {
94 result.tls_1_3 = std.mem.eql(u8, extension.data, &.{ 0x03, 0x04 });
95 },
96 .key_share => result.key_share = try serverKeyShare(extension.data),
97 else => return error.UnsupportedExtension,
98 }
99 }
100 unreachable;
101 }
102
103 pub fn encryptedSelection(
104 extensions: tls.message.Extensions,
105 wanted_alpn: []const u8,
106 server_name_offered: bool,
107 ) Error!EncryptedSelection {
108 var result = EncryptedSelection{};
109 var iterator = extensions.iterator();
110 var seen: [tls.message.extension_count_max]u16 = undefined;
111 var count: u8 = 0;
112 for (0..tls.message.extension_count_max + 1) |_| {
113 const extension = try iterator.next() orelse return result;
114 try mark(&seen, &count, extension.kind);
115 switch (extension.kind) {
116 .application_layer_protocol_negotiation => {
117 result.alpn = try selectedAlpn(extension.data, wanted_alpn);
118 },
119 .quic_transport_parameters => result.transport_parameters = extension.data,
120 .server_certificate_type => result.server_raw_key = try selectedRawKey(extension.data),
121 .client_certificate_type => result.client_raw_key = try selectedRawKey(extension.data),
122 .supported_groups => {},
123 .server_name => {
124 if (!server_name_offered) return error.UnsupportedExtension;
125 if (extension.data.len != 0) return error.InvalidExtension;
126 },
127 else => return error.UnsupportedExtension,
128 }
129 }
130 unreachable;
131 }
132
133 pub fn requestOffersEd25519(extensions: tls.message.Extensions) Error!bool {
134 var iterator = extensions.iterator();
135 var seen: [tls.message.extension_count_max]u16 = undefined;
136 var count: u8 = 0;
137 var offered = false;
138 for (0..tls.message.extension_count_max + 1) |_| {
139 const extension = try iterator.next() orelse return offered;
140 try mark(&seen, &count, extension.kind);
141 switch (extension.kind) {
142 .signature_algorithms => offered = try hasU16(
143 extension.data,
144 @backingInt(std.crypto.tls.SignatureScheme.ed25519),
145 false,
146 ),
147 else => if (known(extension.kind) and
148 !certificateRequestAllowed(extension.kind))
149 {
150 return error.IllegalExtension;
151 },
152 }
153 }
154 unreachable;
155 }
156
157 fn mark(seen: *[tls.message.extension_count_max]u16, count: *u8, kind: ExtensionType) Error!void {
158 const value = @as(u16, @backingInt(kind));
159 for (0..tls.message.extension_count_max) |index| {
160 if (index >= count.*) break;
161 if (seen[index] == value) return error.DuplicateExtension;
162 }
163 std.debug.assert(count.* < tls.message.extension_count_max);
164 seen[count.*] = value;
165 count.* += 1;
166 }
167
168 fn hasU16(data: []const u8, wanted: u16, byte_length: bool) Error!bool {
169 var input = cursor.Read.init(data);
170 const length = if (byte_length) try input.byte() else try input.int(u16);
171 if (length != input.remaining()) return error.InvalidExtension;
172 if (length % 2 != 0) return error.InvalidExtension;
173 var found = false;
174 for (0..profile_list_entries_max) |index| {
175 if (index >= length / 2) break;
176 if (try input.int(u16) == wanted) found = true;
177 }
178 if (input.remaining() != 0) return error.InvalidExtension;
179 return found;
180 }
181
182 fn clientKeyShare(data: []const u8) Error!?[32]u8 {
183 var input = cursor.Read.init(data);
184 const list = try input.take(try input.int(u16));
185 if (input.remaining() != 0) return error.InvalidExtension;
186 var entries = cursor.Read.init(list);
187 var found: ?[32]u8 = null;
188 for (0..profile_list_entries_max) |_| {
189 if (entries.remaining() == 0) return found;
190 const group = try entries.int(u16);
191 const key = try entries.take(try entries.int(u16));
192 if (group != @backingInt(std.crypto.tls.NamedGroup.x25519)) continue;
193 if (key.len != 32) return error.InvalidExtension;
194 if (found != null) return error.InvalidExtension;
195 found = key[0..32].*;
196 }
197 return error.InvalidExtension;
198 }
199
200 fn serverKeyShare(data: []const u8) Error!?[32]u8 {
201 var input = cursor.Read.init(data);
202 const group = try input.int(u16);
203 const key = try input.take(try input.int(u16));
204 if (input.remaining() != 0) return error.InvalidExtension;
205 if (group != @backingInt(std.crypto.tls.NamedGroup.x25519)) return null;
206 if (key.len != 32) return error.InvalidExtension;
207 return key[0..32].*;
208 }
209
210 fn selectAlpn(data: []const u8, wanted: []const u8) Error!?[]const u8 {
211 var input = cursor.Read.init(data);
212 const list = try input.take(try input.int(u16));
213 if (input.remaining() != 0) return error.InvalidExtension;
214 var protocols = cursor.Read.init(list);
215 var found: ?[]const u8 = null;
216 for (0..profile_list_entries_max) |_| {
217 if (protocols.remaining() == 0) return found;
218 const protocol = try protocols.take(try protocols.byte());
219 if (protocol.len == 0) return error.InvalidExtension;
220 if (std.mem.eql(u8, protocol, wanted)) found = protocol;
221 }
222 return error.InvalidExtension;
223 }
224
225 fn selectedAlpn(data: []const u8, wanted: []const u8) Error!?[]const u8 {
226 const selected = try selectAlpn(data, wanted) orelse return null;
227 var input = cursor.Read.init(data);
228 const list = try input.take(try input.int(u16));
229 var protocols = cursor.Read.init(list);
230 _ = try protocols.take(try protocols.byte());
231 if (protocols.remaining() != 0) return error.InvalidExtension;
232 return selected;
233 }
234
235 fn offersRawKey(data: []const u8) Error!bool {
236 var input = cursor.Read.init(data);
237 const list = try input.take(try input.byte());
238 if (input.remaining() != 0) return error.InvalidExtension;
239 for (0..profile_list_entries_max) |index| {
240 if (index >= list.len) break;
241 if (list[index] == @backingInt(std.crypto.tls.CertificateType.RawPublicKey)) return true;
242 }
243 return false;
244 }
245
246 fn selectedRawKey(data: []const u8) Error!bool {
247 if (data.len != 1) return error.InvalidExtension;
248 return data[0] == @backingInt(std.crypto.tls.CertificateType.RawPublicKey);
249 }
250
251 fn serverName(data: []const u8) Error!?[]const u8 {
252 var input = cursor.Read.init(data);
253 const list = try input.take(try input.int(u16));
254 if (input.remaining() != 0) return error.InvalidExtension;
255 var names = cursor.Read.init(list);
256 var host: ?[]const u8 = null;
257 for (0..profile_list_entries_max) |_| {
258 if (names.remaining() == 0) return host;
259 const kind = try names.byte();
260 const name = try names.take(try names.int(u16));
261 if (kind != 0) continue;
262 if (host != null) return error.InvalidExtension;
263 if (name.len == 0) return error.InvalidExtension;
264 host = name;
265 }
266 return error.InvalidExtension;
267 }
268
269 fn known(kind: ExtensionType) bool {
270 return switch (kind) {
271 .server_name,
272 .max_fragment_length,
273 .status_request,
274 .supported_groups,
275 .signature_algorithms,
276 .use_srtp,
277 .heartbeat,
278 .application_layer_protocol_negotiation,
279 .signed_certificate_timestamp,
280 .client_certificate_type,
281 .server_certificate_type,
282 .padding,
283 .pre_shared_key,
284 .early_data,
285 .supported_versions,
286 .cookie,
287 .psk_key_exchange_modes,
288 .certificate_authorities,
289 .oid_filters,
290 .post_handshake_auth,
291 .signature_algorithms_cert,
292 .key_share,
293 .quic_transport_parameters,
294 => true,
295 else => false,
296 };
297 }
298
299 fn clientHelloAllowed(kind: ExtensionType) bool {
300 return switch (kind) {
301 .server_name,
302 .max_fragment_length,
303 .status_request,
304 .supported_groups,
305 .signature_algorithms,
306 .use_srtp,
307 .heartbeat,
308 .application_layer_protocol_negotiation,
309 .signed_certificate_timestamp,
310 .client_certificate_type,
311 .server_certificate_type,
312 .padding,
313 .pre_shared_key,
314 .early_data,
315 .supported_versions,
316 .cookie,
317 .psk_key_exchange_modes,
318 .certificate_authorities,
319 .post_handshake_auth,
320 .signature_algorithms_cert,
321 .key_share,
322 .quic_transport_parameters,
323 => true,
324 else => false,
325 };
326 }
327
328 fn certificateRequestAllowed(kind: ExtensionType) bool {
329 return switch (kind) {
330 .status_request,
331 .signature_algorithms,
332 .signed_certificate_timestamp,
333 .certificate_authorities,
334 .oid_filters,
335 .signature_algorithms_cert,
336 => true,
337 else => false,
338 };
339 }
340
341 fn hexBytes(comptime value: []const u8) [value.len / 2]u8 {
342 var result: [value.len / 2]u8 = undefined;
343 _ = std.fmt.hexToBytes(&result, value) catch unreachable;
344 return result;
345 }
346
347 test "RFC 9001 Appendix A.2 ClientHello reaches clientOffer successfully" {
348 const bytes = hexBytes(
349 "010000ed0303ebf8fa56f12939b9584a3896472ec40bb863cfd3e86804fe3a47" ++
350 "f06a2b69484c00000413011302010000c000000010000e00000b6578616d706c" ++
351 "652e636f6dff01000100000a00080006001d0017001800100007000504616c70" ++
352 "6e000500050100000000003300260024001d00209370b2c9caa47fbabaf4559f" ++
353 "edba753de171fa71f50f1ce15d43e994ec74d748002b0003020304000d001000" ++
354 "0e0403050306030203080408050806002d00020101001c000240010039003204" ++
355 "08ffffffffffffffff05048000ffff07048000ffff080110010480007530090110" ++
356 "0f088394c8f03e51570806048000ffff",
357 );
358 const hello = try tls.message.decodeClientHello(&bytes);
359 const offer = try clientOffer(hello.extensions, "alpn");
360 try std.testing.expect(offer.tls_1_3);
361 try std.testing.expect(offer.x25519_group);
362 try std.testing.expect(offer.key_share != null);
363 try std.testing.expectEqualStrings("alpn", offer.alpn.?);
364 try std.testing.expectEqual(@as(usize, 50), offer.transport_parameters.?.len);
365 }
366
367 test "RFC 8446 section 4.2 ignored legal ClientHello extensions remain unhonored" {
368 const bytes = [_]u8{
369 0x00, 0x2a, 0x00, 0x00,
370 0x00, 0x2d, 0x00, 0x02,
371 0x01, 0x01, 0x00, 0x29,
372 0x00, 0x00,
373 };
374 const offer = try clientOffer(.{ .bytes = &bytes }, "tiny/1");
375 try std.testing.expectEqual(@as(?[]const u8, null), offer.alpn);
376 try std.testing.expectEqual(@as(?[32]u8, null), offer.key_share);
377 }
378
379 test "RFC 8446 section 4.2 legal CertificateRequest extensions are ignored" {
380 const bytes = [_]u8{
381 0x00, 0x05, 0x00, 0x00,
382 0x00, 0x2f, 0x00, 0x00,
383 0x00, 0x30, 0x00, 0x00,
384 0x00, 0x0d, 0x00, 0x04,
385 0x00, 0x02, 0x08, 0x07,
386 };
387 try std.testing.expect(try requestOffersEd25519(.{ .bytes = &bytes }));
388 }
389
390 test "RFC 8446 section 4.2 wrong-message ClientHello extension aborts" {
391 const bytes = [_]u8{ 0x00, 0x30, 0x00, 0x00 };
392 try std.testing.expectError(
393 error.IllegalExtension,
394 clientOffer(.{ .bytes = &bytes }, "tiny/1"),
395 );
396 }
397
398 test "RFC 8446 section 4.2 offered EncryptedExtensions values are accepted" {
399 const supported_groups = [_]u8{ 0x00, 0x0a, 0x00, 0x00 };
400 _ = try encryptedSelection(.{ .bytes = &supported_groups }, "tiny/1", false);
401 const server_name = [_]u8{ 0x00, 0x00, 0x00, 0x00 };
402 try std.testing.expectError(
403 error.UnsupportedExtension,
404 encryptedSelection(.{ .bytes = &server_name }, "tiny/1", false),
405 );
406 _ = try encryptedSelection(.{ .bytes = &server_name }, "tiny/1", true);
407 }