authorgravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2022-12-13 20:15:41-07:00
committergravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2023-01-02 16:57:15-07:00
logd2f5d0b1990a160aa1d648531ea5b1df7b2acdce
treeca92f3233708feb032d3fc3f667efcca7d4a296b
parentba44513c2fe363b55b2c534be98179286b832b7e

std.crypto.Tls: parse the ServerHello handshake


3 files changed, 138 insertions(+), 15 deletions(-)

lib/std/crypto/Tls.zig+114-13
......@@ -8,6 +8,13 @@ const assert = std.debug.assert;
88state: State = .start,
99x25519_priv_key: [32]u8 = undefined,
1010x25519_pub_key: [32]u8 = undefined,
11x25519_server_pub_key: [32]u8 = undefined,
12
13const ProtocolVersion = enum(u16) {
14 tls_1_2 = 0x0303,
15 tls_1_3 = 0x0304,
16 _,
17};
1118
1219const State = enum {
1320 /// In this state, all fields are undefined except state.
......@@ -186,6 +193,18 @@ const NamedGroup = enum(u16) {
186193// * length: u24
187194// * data: opaque
188195
196// ServerHello:
197// * ProtocolVersion legacy_version = 0x0303;
198// * Random random;
199// * opaque legacy_session_id_echo<0..32>;
200// * CipherSuite cipher_suite;
201// * uint8 legacy_compression_method = 0;
202// * Extension extensions<6..2^16-1>;
203
204// Extension:
205// * ExtensionType extension_type;
206// * opaque extension_data<0..2^16-1>;
207
189208const CipherSuite = enum(u16) {
190209 TLS_AES_128_GCM_SHA256 = 0x1301,
191210 TLS_AES_256_GCM_SHA384 = 0x1302,
......@@ -259,10 +278,10 @@ pub fn init(tls: *Tls, stream: net.Stream, host: []const u8) !void {
259278
260279 // Extension: key_share
261280 0, 51, // ExtensionType.key_share
262 0x00, 38, // byte length of this extension payload
263 0x00, 36, // byte length of client_shares
281 0, 38, // byte length of this extension payload
282 0, 36, // byte length of client_shares
264283 0x00, 0x1D, // NamedGroup.x25519
265 0x00, 32, // byte length of key_exchange
284 0, 32, // byte length of key_exchange
266285 } ++ tls.x25519_pub_key ++ [_]u8{
267286
268287 // Extension: server_name
......@@ -313,21 +332,103 @@ pub fn init(tls: *Tls, stream: net.Stream, host: []const u8) !void {
313332 try stream.writevAll(&iovecs);
314333
315334 {
316 var buf: [1000]u8 = undefined;
317 const amt = try stream.read(&buf);
318 const resp = buf[0..amt];
319 const ct = @intToEnum(ContentType, resp[0]);
335 var handshake_buf: [4000]u8 = undefined;
336 const plaintext = handshake_buf[0..5];
337 const amt = try stream.readAtLeast(&handshake_buf, plaintext.len);
338 if (amt < plaintext.len) return error.EndOfStream;
339 const ct = @intToEnum(ContentType, plaintext[0]);
340 const frag_len = mem.readIntBig(u16, plaintext[3..][0..2]);
341 const end = plaintext.len + frag_len;
342 if (end > handshake_buf.len) return error.TlsServerHelloTooBig;
343 if (amt < end) {
344 const amt2 = try stream.readAll(handshake_buf[amt..end]);
345 if (amt2 < plaintext.len) return error.EndOfStream;
346 }
347 const frag = handshake_buf[plaintext.len..end];
348
320349 if (ct == .alert) {
321 //const prot_ver = @bitCast(u16, resp[1..][0..2].*);
322 const len = std.mem.readIntBig(u16, resp[3..][0..2]);
323 const alert = resp[5..][0..len];
324 const level = @intToEnum(AlertLevel, alert[0]);
325 const desc = @intToEnum(AlertDescription, alert[1]);
350 const level = @intToEnum(AlertLevel, frag[0]);
351 const desc = @intToEnum(AlertDescription, frag[1]);
326352 std.debug.print("alert: {s} {s}\n", .{ @tagName(level), @tagName(desc) });
327353 std.process.exit(1);
354 } else if (ct == .handshake) {
355 if (frag[0] != @enumToInt(HandshakeType.server_hello)) {
356 return error.TlsUnexpectedMessage;
357 }
358 const length = mem.readIntBig(u24, frag[1..4]);
359 if (4 + length != frag.len) return error.TlsBadLength;
360 const hello = frag[4..];
361 const legacy_version = mem.readIntBig(u16, hello[0..2]);
362 const random = hello[2..34].*;
363 _ = random;
364 const legacy_session_id_echo_len = hello[34];
365 if (legacy_session_id_echo_len != 0) return error.TlsIllegalParameter;
366 const cipher_suite_int = mem.readIntBig(u16, hello[35..37]);
367 const cipher_suite = std.meta.intToEnum(CipherSuite, cipher_suite_int) catch
368 return error.TlsIllegalParameter;
369 std.debug.print("server wants cipher suite {s}\n", .{@tagName(cipher_suite)});
370 const legacy_compression_method = hello[37];
371 _ = legacy_compression_method;
372 const extensions_size = mem.readIntBig(u16, hello[38..40]);
373 if (40 + extensions_size != hello.len) return error.TlsBadLength;
374 var i: usize = 40;
375 var supported_version: u16 = 0;
376 var have_server_pub_key = false;
377 while (i < hello.len) {
378 const et = mem.readIntBig(u16, hello[i..][0..2]);
379 i += 2;
380 const ext_size = mem.readIntBig(u16, hello[i..][0..2]);
381 i += 2;
382 const next_i = i + ext_size;
383 if (next_i > hello.len) return error.TlsBadLength;
384 switch (et) {
385 @enumToInt(ExtensionType.supported_versions) => {
386 if (supported_version != 0) return error.TlsIllegalParameter;
387 supported_version = mem.readIntBig(u16, hello[i..][0..2]);
388 },
389 @enumToInt(ExtensionType.key_share) => {
390 if (have_server_pub_key) return error.TlsIllegalParameter;
391 const named_group = mem.readIntBig(u16, hello[i..][0..2]);
392 i += 2;
393 switch (named_group) {
394 @enumToInt(NamedGroup.x25519) => {
395 const key_size = mem.readIntBig(u16, hello[i..][0..2]);
396 i += 2;
397 if (key_size != 32) return error.TlsBadLength;
398 const encrypted_key = hello[i..][0..32].*;
399 const server_pub_key = try crypto.dh.X25519.scalarmult(
400 tls.x25519_priv_key,
401 encrypted_key,
402 );
403 tls.x25519_server_pub_key = server_pub_key;
404 have_server_pub_key = true;
405 },
406 else => {
407 std.debug.print("named group: {x}\n", .{named_group});
408 return error.TlsIllegalParameter;
409 },
410 }
411 },
412 else => {
413 std.debug.print("unexpected extension: {x}\n", .{et});
414 },
415 }
416 i = next_i;
417 }
418 if (!have_server_pub_key) return error.TlsIllegalParameter;
419 const tls_version = if (supported_version == 0) legacy_version else supported_version;
420 switch (tls_version) {
421 @enumToInt(ProtocolVersion.tls_1_2) => {
422 std.debug.print("server wants TLS v1.2\n", .{});
423 },
424 @enumToInt(ProtocolVersion.tls_1_3) => {
425 std.debug.print("server wants TLS v1.3\n", .{});
426 },
427 else => return error.TlsIllegalParameter,
428 }
328429 } else {
329430 std.debug.print("content_type: {s}\n", .{@tagName(ct)});
330 std.debug.print("got {d} bytes: {s}\n", .{ amt, std.fmt.fmtSliceHexLower(resp) });
431 std.debug.print("got {d} bytes: {s}\n", .{ amt, std.fmt.fmtSliceHexLower(frag) });
331432 }
332433 }
333434
lib/std/http/Client.zig+2-2
......@@ -59,7 +59,7 @@ pub const Request = struct {
5959
6060pub fn deinit(client: *Client) void {
6161 assert(client.active_requests == 0);
62 client.headers.denit(client.allocator);
62 client.headers.deinit(client.allocator);
6363 client.* = undefined;
6464}
6565
......@@ -69,6 +69,7 @@ pub fn request(client: *Client, options: Request.Options) !Request {
6969 .stream = try net.tcpConnectToHost(client.allocator, options.host, options.port),
7070 .protocol = options.protocol,
7171 };
72 client.active_requests += 1;
7273 errdefer req.deinit();
7374
7475 switch (options.protocol) {
......@@ -100,7 +101,6 @@ pub fn request(client: *Client, options: Request.Options) !Request {
100101 }
101102 req.headers.appendSliceAssumeCapacity(client.headers.items);
102103
103 client.active_requests += 1;
104104 return req;
105105}
106106
lib/std/net.zig+22
......@@ -1672,6 +1672,28 @@ pub const Stream = struct {
16721672 }
16731673 }
16741674
1675 /// Returns the number of bytes read. If the number read is smaller than
1676 /// `buffer.len`, it means the stream reached the end. Reaching the end of
1677 /// a stream is not an error condition.
1678 pub fn readAll(s: Stream, buffer: []u8) ReadError!usize {
1679 return readAtLeast(s, buffer, buffer.len);
1680 }
1681
1682 /// Returns the number of bytes read, calling the underlying read function
1683 /// multiple times until at least the buffer has at least `len` bytes
1684 /// filled. If the number read is less than `len` it means the stream
1685 /// reached the end. Reaching the end of the stream is not an error
1686 /// condition.
1687 pub fn readAtLeast(s: Stream, buffer: []u8, len: usize) ReadError!usize {
1688 var index: usize = 0;
1689 while (index < len) {
1690 const amt = try s.read(buffer[index..]);
1691 if (amt == 0) break;
1692 index += amt;
1693 }
1694 return index;
1695 }
1696
16751697 /// TODO in evented I/O mode, this implementation incorrectly uses the event loop's
16761698 /// file system thread instead of non-blocking. It needs to be reworked to properly
16771699 /// use non-blocking I/O.