lib/quic/src/tls/message/codec.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const quic = @import("../../root.zig");
  3 
  4 const cursor = quic.cursor;
  5 const HandshakeType = std.crypto.tls.HandshakeType;
  6 const ExtensionType = std.crypto.tls.ExtensionType;
  7 
  8 pub const extension_count_max: u8 = 64;
  9 pub const certificate_entries_max: u8 = 8;
 10 
 11 const DecodeSpecific = error{
 12     DuplicateExtension,
 13     InvalidLength,
 14     TooManyCertificateEntries,
 15     TooManyExtensions,
 16     WrongMessage,
 17 };
 18 
 19 const EncodeSpecific = error{InvalidLength};
 20 
 21 pub const DecodeError = cursor.ReadError || DecodeSpecific;
 22 pub const EncodeError = cursor.WriteError || EncodeSpecific;
 23 
 24 pub const Handshake = struct {
 25     kind: HandshakeType,
 26     body: []const u8,
 27 };
 28 
 29 pub const Extension = struct {
 30     kind: ExtensionType,
 31     data: []const u8,
 32 };
 33 
 34 pub const Extensions = struct {
 35     bytes: []const u8,
 36 
 37     pub fn iterator(self: Extensions) ExtensionIterator {
 38         return .{ .input = cursor.Read.init(self.bytes), .count = 0 };
 39     }
 40 };
 41 
 42 pub const ExtensionIterator = struct {
 43     input: cursor.Read,
 44     count: u8,
 45 
 46     pub fn next(self: *ExtensionIterator) DecodeError!?Extension {
 47         if (self.input.remaining() == 0) return null;
 48         if (self.count == extension_count_max) return error.TooManyExtensions;
 49         self.count += 1;
 50         const kind: ExtensionType = @fromBackingInt(@intCast(try self.input.int(u16)));
 51         const length = try self.input.int(u16);
 52         return .{ .kind = kind, .data = try self.input.take(length) };
 53     }
 54 };
 55 
 56 pub const ClientHello = struct {
 57     legacy_version: u16,
 58     random: [32]u8,
 59     legacy_session_id: []const u8,
 60     cipher_suites: []const u8,
 61     compression_methods: []const u8,
 62     extensions: Extensions,
 63 };
 64 
 65 pub const ServerHello = struct {
 66     legacy_version: u16,
 67     random: [32]u8,
 68     legacy_session_id_echo: []const u8,
 69     cipher_suite: u16,
 70     compression_method: u8,
 71     extensions: Extensions,
 72 };
 73 
 74 pub const EncryptedExtensions = struct { extensions: Extensions };
 75 
 76 pub const CertificateRequest = struct {
 77     context: []const u8,
 78     extensions: Extensions,
 79 };
 80 
 81 pub const Certificate = struct {
 82     request_context: []const u8,
 83     entries: CertificateEntries,
 84 };
 85 
 86 pub const CertificateEntry = struct {
 87     data: []const u8,
 88     extensions: Extensions,
 89 };
 90 
 91 pub const CertificateEntries = struct {
 92     bytes: []const u8,
 93 
 94     pub fn iterator(self: CertificateEntries) CertificateIterator {
 95         return .{ .input = cursor.Read.init(self.bytes), .count = 0 };
 96     }
 97 };
 98 
 99 pub const CertificateIterator = struct {
100     input: cursor.Read,
101     count: u8,
102 
103     pub fn next(self: *CertificateIterator) DecodeError!?CertificateEntry {
104         if (self.input.remaining() == 0) return null;
105         if (self.count == certificate_entries_max) {
106             return error.TooManyCertificateEntries;
107         }
108         self.count += 1;
109         const data = try takeVector24(&self.input);
110         const extensions = try takeVector16(&self.input);
111         return .{ .data = data, .extensions = .{ .bytes = extensions } };
112     }
113 };
114 
115 pub const CertificateVerify = struct {
116     algorithm: std.crypto.tls.SignatureScheme,
117     signature: []const u8,
118 };
119 
120 pub const Finished = struct { verify_data: []const u8 };
121 
122 pub const NewSessionTicket = struct {
123     lifetime: u32,
124     age_add: u32,
125     nonce: []const u8,
126     ticket: []const u8,
127     extensions: Extensions,
128 };
129 
130 pub fn decode(bytes: []const u8) DecodeError!Handshake {
131     var input = cursor.Read.init(bytes);
132     const kind: HandshakeType = @fromBackingInt(@intCast(try input.byte()));
133     const length = try input.int(u24);
134     const body = try input.take(length);
135     if (input.remaining() != 0) return error.InvalidLength;
136     return .{ .kind = kind, .body = body };
137 }
138 
139 pub fn encode(value: Handshake, output: *cursor.Write) EncodeError!void {
140     try writeHeader(value.kind, value.body.len, output);
141     try output.put(value.body);
142 }
143 
144 pub fn findExtension(
145     extensions: Extensions,
146     kind: ExtensionType,
147 ) DecodeError!?[]const u8 {
148     var iterator = extensions.iterator();
149     var found: ?[]const u8 = null;
150     for (0..extension_count_max + 1) |_| {
151         const extension = try iterator.next() orelse return found;
152         if (@backingInt(extension.kind) != @backingInt(kind)) continue;
153         if (found != null) return error.DuplicateExtension;
154         found = extension.data;
155     }
156     unreachable;
157 }
158 
159 fn expected(bytes: []const u8, kind: HandshakeType) DecodeError!cursor.Read {
160     const value = try decode(bytes);
161     if (value.kind != kind) return error.WrongMessage;
162     return cursor.Read.init(value.body);
163 }
164 
165 fn takeVector8(input: *cursor.Read) DecodeError![]const u8 {
166     return input.take(try input.byte());
167 }
168 
169 fn takeVector16(input: *cursor.Read) DecodeError![]const u8 {
170     return input.take(try input.int(u16));
171 }
172 
173 fn takeVector24(input: *cursor.Read) DecodeError![]const u8 {
174     return input.take(try input.int(u24));
175 }
176 
177 fn finish(input: *const cursor.Read) DecodeError!void {
178     if (input.remaining() != 0) return error.InvalidLength;
179 }
180 
181 pub fn decodeClientHello(bytes: []const u8) DecodeError!ClientHello {
182     var input = try expected(bytes, .client_hello);
183     const legacy_version = try input.int(u16);
184     var random: [32]u8 = undefined;
185     @memcpy(&random, try input.take(32));
186     const session_id = try takeVector8(&input);
187     const cipher_suites = try takeVector16(&input);
188     if (cipher_suites.len == 0) return error.InvalidLength;
189     if (cipher_suites.len % 2 != 0) return error.InvalidLength;
190     const compression = try takeVector8(&input);
191     const extensions = try takeVector16(&input);
192     try finish(&input);
193     return .{
194         .legacy_version = legacy_version,
195         .random = random,
196         .legacy_session_id = session_id,
197         .cipher_suites = cipher_suites,
198         .compression_methods = compression,
199         .extensions = .{ .bytes = extensions },
200     };
201 }
202 
203 pub fn encodeClientHello(value: ClientHello, output: *cursor.Write) EncodeError!void {
204     if (value.cipher_suites.len == 0) return error.InvalidLength;
205     if (value.cipher_suites.len % 2 != 0) return error.InvalidLength;
206     const length = try helloLength(
207         value.legacy_session_id.len,
208         value.cipher_suites.len,
209         value.compression_methods.len,
210         value.extensions.bytes.len,
211     );
212     try writeHeader(.client_hello, length, output);
213     try output.int(u16, value.legacy_version);
214     try output.put(&value.random);
215     try writeVector8(value.legacy_session_id, output);
216     try writeVector16(value.cipher_suites, output);
217     try writeVector8(value.compression_methods, output);
218     try writeVector16(value.extensions.bytes, output);
219 }
220 
221 fn helloLength(session: usize, suites: usize, compression: usize, extensions: usize) !usize {
222     var length: usize = 2 + 32 + 1 + 2 + 1 + 2;
223     length = try addLength(length, session);
224     length = try addLength(length, suites);
225     length = try addLength(length, compression);
226     return addLength(length, extensions);
227 }
228 
229 pub fn decodeServerHello(bytes: []const u8) DecodeError!ServerHello {
230     var input = try expected(bytes, .server_hello);
231     const legacy_version = try input.int(u16);
232     var random: [32]u8 = undefined;
233     @memcpy(&random, try input.take(32));
234     const session_id = try takeVector8(&input);
235     const cipher_suite = try input.int(u16);
236     const compression_method = try input.byte();
237     const extensions = try takeVector16(&input);
238     try finish(&input);
239     return .{
240         .legacy_version = legacy_version,
241         .random = random,
242         .legacy_session_id_echo = session_id,
243         .cipher_suite = cipher_suite,
244         .compression_method = compression_method,
245         .extensions = .{ .bytes = extensions },
246     };
247 }
248 
249 pub fn encodeServerHello(value: ServerHello, output: *cursor.Write) EncodeError!void {
250     var length: usize = 2 + 32 + 1 + 2 + 1 + 2;
251     length = try addLength(length, value.legacy_session_id_echo.len);
252     length = try addLength(length, value.extensions.bytes.len);
253     try writeHeader(.server_hello, length, output);
254     try output.int(u16, value.legacy_version);
255     try output.put(&value.random);
256     try writeVector8(value.legacy_session_id_echo, output);
257     try output.int(u16, value.cipher_suite);
258     try output.byte(value.compression_method);
259     try writeVector16(value.extensions.bytes, output);
260 }
261 
262 pub fn decodeEncryptedExtensions(bytes: []const u8) DecodeError!EncryptedExtensions {
263     var input = try expected(bytes, .encrypted_extensions);
264     const extensions = try takeVector16(&input);
265     try finish(&input);
266     return .{ .extensions = .{ .bytes = extensions } };
267 }
268 
269 pub fn encodeEncryptedExtensions(
270     value: EncryptedExtensions,
271     output: *cursor.Write,
272 ) EncodeError!void {
273     const length = try addLength(2, value.extensions.bytes.len);
274     try writeHeader(.encrypted_extensions, length, output);
275     try writeVector16(value.extensions.bytes, output);
276 }
277 
278 pub fn decodeCertificateRequest(bytes: []const u8) DecodeError!CertificateRequest {
279     var input = try expected(bytes, .certificate_request);
280     const context = try takeVector8(&input);
281     const extensions = try takeVector16(&input);
282     try finish(&input);
283     return .{ .context = context, .extensions = .{ .bytes = extensions } };
284 }
285 
286 pub fn encodeCertificateRequest(
287     value: CertificateRequest,
288     output: *cursor.Write,
289 ) EncodeError!void {
290     var length = try addLength(1, value.context.len);
291     length = try addLength(length, 2);
292     length = try addLength(length, value.extensions.bytes.len);
293     try writeHeader(.certificate_request, length, output);
294     try writeVector8(value.context, output);
295     try writeVector16(value.extensions.bytes, output);
296 }
297 
298 pub fn decodeCertificate(bytes: []const u8) DecodeError!Certificate {
299     var input = try expected(bytes, .certificate);
300     const context = try takeVector8(&input);
301     const entries = try takeVector24(&input);
302     try finish(&input);
303     return .{ .request_context = context, .entries = .{ .bytes = entries } };
304 }
305 
306 pub fn encodeCertificate(value: Certificate, output: *cursor.Write) EncodeError!void {
307     var length = try addLength(1, value.request_context.len);
308     length = try addLength(length, 3);
309     length = try addLength(length, value.entries.bytes.len);
310     try writeHeader(.certificate, length, output);
311     try writeVector8(value.request_context, output);
312     try writeVector24(value.entries.bytes, output);
313 }
314 
315 pub fn encodeSingleCertificate(
316     request_context: []const u8,
317     certificate: []const u8,
318     output: *cursor.Write,
319 ) EncodeError!void {
320     var entries_length = try addLength(3, certificate.len);
321     entries_length = try addLength(entries_length, 2);
322     var length = try addLength(1, request_context.len);
323     length = try addLength(length, 3);
324     length = try addLength(length, entries_length);
325     try writeHeader(.certificate, length, output);
326     try writeVector8(request_context, output);
327     try writeLength(u24, entries_length, output);
328     try writeVector24(certificate, output);
329     try writeVector16(&.{}, output);
330 }
331 
332 pub fn decodeCertificateVerify(bytes: []const u8) DecodeError!CertificateVerify {
333     var input = try expected(bytes, .certificate_verify);
334     const algorithm: std.crypto.tls.SignatureScheme = @fromBackingInt(@intCast(try input.int(u16)));
335     const signature = try takeVector16(&input);
336     try finish(&input);
337     return .{ .algorithm = algorithm, .signature = signature };
338 }
339 
340 pub fn encodeCertificateVerify(
341     value: CertificateVerify,
342     output: *cursor.Write,
343 ) EncodeError!void {
344     const length = try addLength(4, value.signature.len);
345     try writeHeader(.certificate_verify, length, output);
346     try output.int(u16, @backingInt(value.algorithm));
347     try writeVector16(value.signature, output);
348 }
349 
350 pub fn decodeNewSessionTicket(bytes: []const u8) DecodeError!NewSessionTicket {
351     var input = try expected(bytes, .new_session_ticket);
352     const lifetime = try input.int(u32);
353     const age_add = try input.int(u32);
354     const nonce = try takeVector8(&input);
355     const ticket = try takeVector16(&input);
356     if (ticket.len == 0) return error.InvalidLength;
357     const extensions = try takeVector16(&input);
358     if (extensions.len > std.math.maxInt(u16) - 1) return error.InvalidLength;
359     try finish(&input);
360     return .{
361         .lifetime = lifetime,
362         .age_add = age_add,
363         .nonce = nonce,
364         .ticket = ticket,
365         .extensions = .{ .bytes = extensions },
366     };
367 }
368 
369 pub fn decodeFinished(bytes: []const u8) DecodeError!Finished {
370     const value = try decode(bytes);
371     if (value.kind != .finished) return error.WrongMessage;
372     return .{ .verify_data = value.body };
373 }
374 
375 pub fn encodeFinished(value: Finished, output: *cursor.Write) EncodeError!void {
376     try writeHeader(.finished, value.verify_data.len, output);
377     try output.put(value.verify_data);
378 }
379 
380 fn writeHeader(kind: HandshakeType, length: usize, output: *cursor.Write) EncodeError!void {
381     try output.byte(@backingInt(kind));
382     try writeLength(u24, length, output);
383 }
384 
385 fn writeVector8(bytes: []const u8, output: *cursor.Write) EncodeError!void {
386     try writeLength(u8, bytes.len, output);
387     try output.put(bytes);
388 }
389 
390 fn writeVector16(bytes: []const u8, output: *cursor.Write) EncodeError!void {
391     try writeLength(u16, bytes.len, output);
392     try output.put(bytes);
393 }
394 
395 fn writeVector24(bytes: []const u8, output: *cursor.Write) EncodeError!void {
396     try writeLength(u24, bytes.len, output);
397     try output.put(bytes);
398 }
399 
400 fn writeLength(comptime T: type, length: usize, output: *cursor.Write) EncodeError!void {
401     if (length > std.math.maxInt(T)) return error.InvalidLength;
402     try output.int(T, @intCast(length));
403 }
404 
405 fn addLength(left: usize, right: usize) EncodeError!usize {
406     return std.math.add(usize, left, right) catch error.InvalidLength;
407 }