| ... | ... | @@ -188,6 +188,12 @@ const NamedGroup = enum(u16) { |
| 188 | 188 | // * fragment: opaque |
| 189 | 189 | // - the data being transmitted |
| 190 | 190 | |
| 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 | |
| 191 | 197 | // Handshake: |
| 192 | 198 | // * type: HandshakeType |
| 193 | 199 | // * length: u24 |
| ... | ... | @@ -331,105 +337,144 @@ pub fn init(tls: *Tls, stream: net.Stream, host: []const u8) !void { |
| 331 | 337 | }; |
| 332 | 338 | try stream.writevAll(&iovecs); |
| 333 | 339 | |
| 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: { |
| 336 | 343 | 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; |
| 339 | 346 | const ct = @intToEnum(ContentType, plaintext[0]); |
| 340 | 347 | const frag_len = mem.readIntBig(u16, plaintext[3..][0..2]); |
| 341 | 348 | 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; |
| 346 | 353 | } |
| 347 | 354 | const frag = handshake_buf[plaintext.len..end]; |
| 348 | 355 | |
| 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", .{}); |
| 411 | 432 | }, |
| 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", .{}); |
| 414 | 435 | }, |
| 436 | else => return error.TlsIllegalParameter, |
| 415 | 437 | } |
| 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 | }, |
| 432 | 476 | } |
| 477 | i = end; |
| 433 | 478 | } |
| 434 | 479 | |
| 435 | 480 | tls.state = .sent_hello; |