Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
254 changes: 247 additions & 7 deletions src/node.zig
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
//! Inbound UDP handling: session lookup, WHOAREYOU challenges, HANDSHAKE completion, and encrypted PING/PONG.
//! Discovery v5 node: inbound `handleReceive`, outbound handshake opening, session cache, and encrypted PING/PONG.

const std = @import("std");
const builtin = @import("builtin");
Expand All @@ -25,18 +25,26 @@ pub const Config = struct {
};

pub const Node = struct {
const PendingChallenge = struct {
peer_id: NodeId,
challenge_data: []u8,
};

const OutboundHandshake = struct {
peer_id: NodeId,
peer_pubkey_compressed: [33]u8,
remote: RemoteUdp,
message_nonce: [12]u8,
};

allocator: std.mem.Allocator,
secret_key: [32]u8,
node_id: NodeId,
enr_seq: u64,
sessions: session.SessionTable,
routing: routing.RoutingTable,
pending: std.ArrayList(PendingChallenge),

const PendingChallenge = struct {
peer_id: NodeId,
challenge_data: []u8,
};
outbound: std.ArrayList(OutboundHandshake),

pub const InitError = identity_v4.EcdhError || session.InitError;

Expand All @@ -52,12 +60,14 @@ pub const Node = struct {
.sessions = sessions,
.routing = routing.RoutingTable.init(nid),
.pending = .empty,
.outbound = .empty,
};
}

pub fn deinit(self: *Node) void {
for (self.pending.items) |p| self.allocator.free(p.challenge_data);
self.pending.deinit(self.allocator);
self.outbound.deinit(self.allocator);
self.sessions.deinit();
}

Expand Down Expand Up @@ -88,17 +98,94 @@ pub const Node = struct {
return null;
}

pub const OpeningHandshakeError = packet.EncodeError ||
message_crypto.Error ||
std.mem.Allocator.Error ||
identity_v4.EcdhError ||
error{PeerKeyIdMismatch};

pub const ReceiveError = packet.Error ||
message_crypto.Error ||
message.DecodeError ||
enr.Error ||
identity_v4.IdentityProofVerifyError ||
identity_v4.IdentityProofSignError ||
identity_v4.EcdhError ||
packet.EncodeError ||
routing.Error ||
std.mem.Allocator.Error ||
error{ MissingHandshakePending, EmptyHandshakeRecord, EnrNodeIdMismatch, BadHandshakeSignatureLength };

fn clearOutboundForPeer(self: *Node, peer_id: NodeId) void {
var i: usize = 0;
while (i < self.outbound.items.len) {
if (std.mem.eql(u8, &self.outbound.items[i].peer_id, &peer_id)) {
_ = self.outbound.swapRemove(i);
} else i += 1;
}
}

fn takeOutboundByMessageNonce(self: *Node, message_nonce: [12]u8) ?OutboundHandshake {
var i: usize = 0;
while (i < self.outbound.items.len) {
if (std.mem.eql(u8, &self.outbound.items[i].message_nonce, &message_nonce)) {
return self.outbound.swapRemove(i);
}
i += 1;
}
return null;
}

/// First encrypted ordinary (unknown session) toward `peer_id`, using a random throwaway AES key.
/// Completes after the peer's WHOAREYOU is passed to **handleReceive**. Caller frees the returned datagram.
pub fn allocOpeningPingHandshake(
self: *Node,
peer_id: NodeId,
peer_pubkey_compressed: [33]u8,
remote: RemoteUdp,
req_id: []const u8,
ping_enr_seq: u64,
) OpeningHandshakeError![]u8 {
const alloc = self.allocator;
const derived = try identity_v4.nodeIdV4FromCompressedSec1(peer_pubkey_compressed);
if (!std.mem.eql(u8, &derived, &peer_id)) return error.PeerKeyIdMismatch;

self.clearOutboundForPeer(peer_id);

var iv: [16]u8 = undefined;
fillRandomBytes(&iv);
var nonce: [12]u8 = undefined;
fillRandomBytes(&nonce);

var temp_key: [16]u8 = undefined;
fillRandomBytes(&temp_key);
defer @memset(&temp_key, 0);

const ping_pt = try message.encodePingPlaintext(alloc, req_id, ping_enr_seq);
defer alloc.free(ping_pt);

var prefix: [packet.static_prefix_size + packet.message_auth_size]u8 = undefined;
@memcpy(prefix[0..16], &iv);
var static_plain: [packet.static_header_size]u8 = undefined;
packet.writePlaintextStaticHeader(&static_plain, .message, nonce, packet.message_auth_size);
@memcpy(prefix[16..][0..packet.static_header_size], &static_plain);
@memcpy(prefix[packet.static_prefix_size..][0..packet.message_auth_size], &self.node_id);

const ct = try message_crypto.encryptMessage(alloc, temp_key, nonce, ping_pt, &prefix);
defer alloc.free(ct);

const pkt = try packet.encodeOrdinaryMessagePacket(alloc, peer_id, iv, nonce, self.node_id, ct);
errdefer alloc.free(pkt);

try self.outbound.append(alloc, .{
.peer_id = peer_id,
.peer_pubkey_compressed = peer_pubkey_compressed,
.remote = remote,
.message_nonce = nonce,
});
return pkt;
}

/// On success, each element of `responses_out` is an allocated reply packet; caller must free them.
pub fn handleReceive(
self: *Node,
Expand All @@ -113,12 +200,70 @@ pub const Node = struct {

const parsed = try packet.decodeInPlace(&self.node_id, copy);
switch (parsed.header.flag) {
.whoareyou => return,
.whoareyou => try self.handleWhoareyouAsInitiator(remote, copy, parsed, responses_out),
.message => try self.handleOrdinary(remote, copy, parsed, responses_out),
.handshake => try self.handleHandshake(remote, copy, parsed, responses_out),
}
}

fn handleWhoareyouAsInitiator(
self: *Node,
remote: RemoteUdp,
_: []u8,
parsed: packet.ParsedPacket,
responses_out: *std.ArrayList([]u8),
) ReceiveError!void {
const alloc = self.allocator;
const o = self.takeOutboundByMessageNonce(parsed.header.nonce) orelse return;

_ = try parsed.decodeAuth();

var cd_buf: [packet.static_prefix_size + packet.whoareyou_auth_size]u8 = undefined;
packet.writeWhoareyouChallengeData(&cd_buf, parsed);

var sk_eph: [32]u8 = undefined;
fillRandomBytes(&sk_eph);
defer @memset(&sk_eph, 0);

const eph_pub = try identity_v4.compressedPubkeyFromSecretKey(sk_eph);
const ikm = try identity_v4.ecdhLocalSecret(o.peer_pubkey_compressed, sk_eph);

const keys = handshake.deriveSessionKeys(&ikm, &cd_buf, self.node_id, o.peer_id);
const sig = try identity_v4.signIdentityProof(&cd_buf, &eph_pub, o.peer_id, self.secret_key, null);

const pk_self = try identity_v4.compressedPubkeyFromSecretKey(self.secret_key);
const record = try buildMinimalEnrRlp(alloc, pk_self, self.enr_seq);
defer alloc.free(record);

var iv1: [16]u8 = undefined;
fillRandomBytes(&iv1);
var nonce_hs: [12]u8 = undefined;
fillRandomBytes(&nonce_hs);

const hs_pkt = try packet.encodeHandshakePacket(
alloc,
o.peer_id,
iv1,
nonce_hs,
self.node_id,
64,
33,
&sig,
&eph_pub,
record,
&.{},
);
errdefer alloc.free(hs_pkt);

const ep = self.makeEndpoint(o.peer_id, remote.ip, remote.port);
try self.sessions.put(ep, session.CachedSession.fromDerived(keys), false);
errdefer _ = self.sessions.remove(ep);

_ = try self.routing.add(o.peer_id);

try responses_out.append(alloc, hs_pkt);
}

fn handleOrdinary(
self: *Node,
remote: RemoteUdp,
Expand Down Expand Up @@ -573,3 +718,98 @@ test "responder completes handshake and answers ping inside handshake" {
try std.testing.expect(msg2 == .pong);
try std.testing.expectEqualSlices(u8, &.{0x11}, msg2.pong.req_id);
}

test "initiator opening ping completes handshake after WHOAREYOU" {
const alloc = std.testing.allocator;

var sk_b: [32]u8 = @splat(0);
sk_b[31] = 201;
var node_b = try Node.init(alloc, .{ .secret_key = sk_b, .enr_seq = 4 });
defer node_b.deinit();

var sk_a: [32]u8 = @splat(0);
sk_a[31] = 203;
var node_a = try Node.init(alloc, .{ .secret_key = sk_a, .enr_seq = 2 });
defer node_a.deinit();

const id_a = node_a.node_id;
const id_b = node_b.node_id;
const pk_b = try identity_v4.compressedPubkeyFromSecretKey(sk_b);

const remote_a: RemoteUdp = .{ .ip = .{ .v4 = .{ 10, 0, 0, 11 } }, .port = 51111 };
const remote_b: RemoteUdp = .{ .ip = .{ .v4 = .{ 10, 0, 0, 22 } }, .port = 52222 };

const ordinary = try node_a.allocOpeningPingHandshake(id_b, pk_b, remote_b, &.{ 0xca, 0xfe }, 9);
defer alloc.free(ordinary);

var from_b: std.ArrayList([]u8) = .empty;
defer {
for (from_b.items) |s| alloc.free(s);
from_b.deinit(alloc);
}
try node_b.handleReceive(remote_a, ordinary, &from_b);
try std.testing.expectEqual(@as(usize, 1), from_b.items.len);
const way = from_b.items[0];

var from_a: std.ArrayList([]u8) = .empty;
defer {
for (from_a.items) |s| alloc.free(s);
from_a.deinit(alloc);
}
try node_a.handleReceive(remote_b, way, &from_a);
try std.testing.expectEqual(@as(usize, 1), from_a.items.len);
const hs = from_a.items[0];

for (from_b.items) |s| alloc.free(s);
from_b.clearRetainingCapacity();

try node_b.handleReceive(remote_a, hs, &from_b);
try std.testing.expectEqual(@as(usize, 0), from_b.items.len);

const ep_on_a = node_a.makeEndpoint(id_b, remote_b.ip, remote_b.port);
const ep_on_b = node_b.makeEndpoint(id_a, remote_a.ip, remote_a.port);
const lu_a = node_a.sessions.get(ep_on_a) orelse unreachable;
const lu_b = node_b.sessions.get(ep_on_b) orelse unreachable;
try std.testing.expect(!lu_a.peer_handshake_initiator);
try std.testing.expect(lu_b.peer_handshake_initiator);

var iv: [16]u8 = undefined;
@memset(&iv, 0x5e);
var nonce: [12]u8 = undefined;
@memset(&nonce, 0x6f);

const ping_pt2 = try message.encodePingPlaintext(alloc, &.{ 0x01, 0x02 }, 3);
defer alloc.free(ping_pt2);

var prefix: [packet.static_prefix_size + packet.message_auth_size]u8 = undefined;
@memcpy(prefix[0..16], &iv);
var static_plain: [packet.static_header_size]u8 = undefined;
packet.writePlaintextStaticHeader(&static_plain, .message, nonce, packet.message_auth_size);
@memcpy(prefix[16..][0..packet.static_header_size], &static_plain);
@memcpy(prefix[packet.static_prefix_size..][0..packet.message_auth_size], &id_a);

const write_key = lu_a.session.writeKeyWeWereInitiator();
const ct2 = try message_crypto.encryptMessage(alloc, write_key, nonce, ping_pt2, &prefix);
defer alloc.free(ct2);

const ordinary2 = try packet.encodeOrdinaryMessagePacket(alloc, id_b, iv, nonce, id_a, ct2);
defer alloc.free(ordinary2);

var from_b2: std.ArrayList([]u8) = .empty;
defer {
for (from_b2.items) |s| alloc.free(s);
from_b2.deinit(alloc);
}
try node_b.handleReceive(remote_a, ordinary2, &from_b2);
try std.testing.expectEqual(@as(usize, 1), from_b2.items.len);

const pong_copy = try alloc.dupe(u8, from_b2.items[0]);
defer alloc.free(pong_copy);
const parsed_pong = try packet.decodeInPlace(&id_a, pong_copy);
const read_key = lu_a.session.readKeyWeWereInitiator();
const plain = try message_crypto.decryptOrdinaryMessage(alloc, pong_copy, &parsed_pong, read_key);
defer alloc.free(plain);
const msg2 = try message.decodePlaintext(plain, alloc);
try std.testing.expect(msg2 == .pong);
try std.testing.expectEqualSlices(u8, &.{ 0x01, 0x02 }, msg2.pong.req_id);
}
11 changes: 11 additions & 0 deletions src/packet.zig
Original file line number Diff line number Diff line change
Expand Up @@ -216,6 +216,17 @@ pub fn allocWhoareyouChallengeData(
return out;
}

/// Writes the WHOAREYOU challenge bytes used as HKDF salt in **handshake.deriveSessionKeys** (unmasked prefix).
pub fn writeWhoareyouChallengeData(out: *[static_prefix_size + whoareyou_auth_size]u8, parsed: ParsedPacket) void {
std.debug.assert(parsed.header.flag == .whoareyou);
std.debug.assert(parsed.auth_data.len == whoareyou_auth_size);
@memcpy(out[0..16], &parsed.iv);
var static_plain: [static_header_size]u8 = undefined;
writePlaintextStaticHeader(&static_plain, .whoareyou, parsed.header.nonce, whoareyou_auth_size);
@memcpy(out[16..][0..static_header_size], &static_plain);
@memcpy(out[static_prefix_size..], parsed.auth_data);
}

pub fn encodeWhoareyouPacket(
allocator: std.mem.Allocator,
dest_node_id: [32]u8,
Expand Down
Loading