authorgravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2022-12-13 21:59:01-07:00
committergravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2023-01-02 16:57:15-07:00
log920e5bc4ff4bdfee173768809e712f8004f7132d
treefb9d7a7d8c85ea9cb4b329b2b0fe9bff69859e83
parentd2f5d0b1990a160aa1d648531ea5b1df7b2acdce

std.crypto.Tls: discard ChangeCipherSpec messages

The next step here is to decrypt encrypted records

2 files changed, 136 insertions(+), 91 deletions(-)

lib/std/crypto/Tls.zig+133-88
......@@ -188,6 +188,12 @@ const NamedGroup = enum(u16) {
188188// * fragment: opaque
189189// - the data being transmitted
190190
191// Ciphertext
192// * ContentType opaque_type = application_data; /* 23 */
193// * ProtocolVersion legacy_record_version = 0x0303; /* TLS v1.2 */
194// * uint16 length;
195// * opaque encrypted_record[TLSCiphertext.length];
196
191197// Handshake:
192198// * type: HandshakeType
193199// * length: u24
......@@ -331,105 +337,144 @@ pub fn init(tls: *Tls, stream: net.Stream, host: []const u8) !void {
331337 };
332338 try stream.writevAll(&iovecs);
333339
334 {
335 var handshake_buf: [4000]u8 = undefined;
340 var handshake_buf: [4000]u8 = undefined;
341 var len: usize = 0;
342 var i: usize = i: {
336343 const plaintext = handshake_buf[0..5];
337 const amt = try stream.readAtLeast(&handshake_buf, plaintext.len);
338 if (amt < plaintext.len) return error.EndOfStream;
344 len = try stream.readAtLeast(&handshake_buf, plaintext.len);
345 if (len < plaintext.len) return error.EndOfStream;
339346 const ct = @intToEnum(ContentType, plaintext[0]);
340347 const frag_len = mem.readIntBig(u16, plaintext[3..][0..2]);
341348 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;
349 if (end > handshake_buf.len) return error.TlsRecordOverflow;
350 if (end > len) {
351 len += try stream.readAtLeast(handshake_buf[len..], end - len);
352 if (end > len) return error.EndOfStream;
346353 }
347354 const frag = handshake_buf[plaintext.len..end];
348355
349 if (ct == .alert) {
350 const level = @intToEnum(AlertLevel, frag[0]);
351 const desc = @intToEnum(AlertDescription, frag[1]);
352 std.debug.print("alert: {s} {s}\n", .{ @tagName(level), @tagName(desc) });
353 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 }
356 switch (ct) {
357 .alert => {
358 const level = @intToEnum(AlertLevel, frag[0]);
359 const desc = @intToEnum(AlertDescription, frag[1]);
360 std.debug.print("alert: {s} {s}\n", .{ @tagName(level), @tagName(desc) });
361 return error.TlsAlert;
362 },
363 .handshake => {
364 if (frag[0] != @enumToInt(HandshakeType.server_hello)) {
365 return error.TlsUnexpectedMessage;
366 }
367 const length = mem.readIntBig(u24, frag[1..4]);
368 if (4 + length != frag.len) return error.TlsBadLength;
369 const hello = frag[4..];
370 const legacy_version = mem.readIntBig(u16, hello[0..2]);
371 const random = hello[2..34].*;
372 _ = random;
373 const legacy_session_id_echo_len = hello[34];
374 if (legacy_session_id_echo_len != 0) return error.TlsIllegalParameter;
375 const cipher_suite_int = mem.readIntBig(u16, hello[35..37]);
376 const cipher_suite = std.meta.intToEnum(CipherSuite, cipher_suite_int) catch
377 return error.TlsIllegalParameter;
378 std.debug.print("server wants cipher suite {s}\n", .{@tagName(cipher_suite)});
379 const legacy_compression_method = hello[37];
380 _ = legacy_compression_method;
381 const extensions_size = mem.readIntBig(u16, hello[38..40]);
382 if (40 + extensions_size != hello.len) return error.TlsBadLength;
383 var i: usize = 40;
384 var supported_version: u16 = 0;
385 var have_server_pub_key = false;
386 while (i < hello.len) {
387 const et = mem.readIntBig(u16, hello[i..][0..2]);
388 i += 2;
389 const ext_size = mem.readIntBig(u16, hello[i..][0..2]);
390 i += 2;
391 const next_i = i + ext_size;
392 if (next_i > hello.len) return error.TlsBadLength;
393 switch (et) {
394 @enumToInt(ExtensionType.supported_versions) => {
395 if (supported_version != 0) return error.TlsIllegalParameter;
396 supported_version = mem.readIntBig(u16, hello[i..][0..2]);
397 },
398 @enumToInt(ExtensionType.key_share) => {
399 if (have_server_pub_key) return error.TlsIllegalParameter;
400 const named_group = mem.readIntBig(u16, hello[i..][0..2]);
401 i += 2;
402 switch (named_group) {
403 @enumToInt(NamedGroup.x25519) => {
404 const key_size = mem.readIntBig(u16, hello[i..][0..2]);
405 i += 2;
406 if (key_size != 32) return error.TlsBadLength;
407 const encrypted_key = hello[i..][0..32].*;
408 const server_pub_key = try crypto.dh.X25519.scalarmult(
409 tls.x25519_priv_key,
410 encrypted_key,
411 );
412 tls.x25519_server_pub_key = server_pub_key;
413 have_server_pub_key = true;
414 },
415 else => {
416 std.debug.print("named group: {x}\n", .{named_group});
417 return error.TlsIllegalParameter;
418 },
419 }
420 },
421 else => {
422 std.debug.print("unexpected extension: {x}\n", .{et});
423 },
424 }
425 i = next_i;
426 }
427 if (!have_server_pub_key) return error.TlsIllegalParameter;
428 const tls_version = if (supported_version == 0) legacy_version else supported_version;
429 switch (tls_version) {
430 @enumToInt(ProtocolVersion.tls_1_2) => {
431 std.debug.print("server wants TLS v1.2\n", .{});
411432 },
412 else => {
413 std.debug.print("unexpected extension: {x}\n", .{et});
433 @enumToInt(ProtocolVersion.tls_1_3) => {
434 std.debug.print("server wants TLS v1.3\n", .{});
414435 },
436 else => return error.TlsIllegalParameter,
415437 }
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 }
429 } else {
430 std.debug.print("content_type: {s}\n", .{@tagName(ct)});
431 std.debug.print("got {d} bytes: {s}\n", .{ amt, std.fmt.fmtSliceHexLower(frag) });
438 },
439 else => return error.TlsUnexpectedMessage,
440 }
441 break :i end;
442 };
443
444 while (true) {
445 const end_hdr = i + 5;
446 if (end_hdr > handshake_buf.len) return error.TlsRecordOverflow;
447 if (end_hdr > len) {
448 len += try stream.readAtLeast(handshake_buf[len..], end_hdr - len);
449 if (end_hdr > len) return error.EndOfStream;
450 }
451 const ct = @intToEnum(ContentType, handshake_buf[i]);
452 i += 1;
453 const legacy_version = mem.readIntBig(u16, handshake_buf[i..][0..2]);
454 i += 2;
455 _ = legacy_version;
456 const record_size = mem.readIntBig(u16, handshake_buf[i..][0..2]);
457 i += 2;
458 const end = i + record_size;
459 if (end > handshake_buf.len) return error.TlsRecordOverflow;
460 if (end > len) {
461 len += try stream.readAtLeast(handshake_buf[len..], end - len);
462 if (end > len) return error.EndOfStream;
463 }
464 switch (ct) {
465 .change_cipher_spec => {
466 if (record_size != 1) return error.TlsUnexpectedMessage;
467 if (handshake_buf[i] != 0x01) return error.TlsUnexpectedMessage;
468 },
469 .application_data => {
470 std.debug.print("TODO: decrypt these {d} bytes\n", .{record_size});
471 },
472 else => {
473 std.debug.print("content type: {s}\n", .{@tagName(ct)});
474 return error.TlsUnexpectedMessage;
475 },
432476 }
477 i = end;
433478 }
434479
435480 tls.state = .sent_hello;
lib/std/net.zig+3-3
......@@ -1680,9 +1680,9 @@ pub const Stream = struct {
16801680 }
16811681
16821682 /// 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
1683 /// the minimal number of times until at least the buffer has at least
1684 /// `len` bytes filled. If the number read is less than `len` it means the
1685 /// stream reached the end. Reaching the end of the stream is not an error
16861686 /// condition.
16871687 pub fn readAtLeast(s: Stream, buffer: []u8, len: usize) ReadError!usize {
16881688 var index: usize = 0;