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 }