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 }