From f8872a0cd19738e5dbfd763d33807add1752d662 Mon Sep 17 00:00:00 2001 From: Tu Dinh Ngoc Date: Thu, 20 Jun 2024 13:28:38 +0000 Subject: [PATCH 01/34] tun: use add-with-carry in checksumNoFold() Use parallel summation with native byte order per RFC 1071. add-with-carry operation is used to add 4 words per operation. Byteswap is performed before and after checksumming for compatibility with old `checksumNoFold()`. With this we get a 30-80% speedup in `checksum()` depending on packet sizes. Add unit tests with comparison to a per-word implementation. **Intel(R) Xeon(R) Silver 4210R CPU @ 2.40GHz** | Size | OldTime | NewTime | Speedup | |------|---------|---------|----------| | 64 | 12.64 | 9.183 | 1.376456 | | 128 | 18.52 | 12.72 | 1.455975 | | 256 | 31.01 | 18.13 | 1.710425 | | 512 | 54.46 | 29.03 | 1.87599 | | 1024 | 102 | 52.2 | 1.954023 | | 1500 | 146.8 | 81.36 | 1.804326 | | 2048 | 196.9 | 102.5 | 1.920976 | | 4096 | 389.8 | 200.8 | 1.941235 | | 8192 | 767.3 | 413.3 | 1.856521 | | 9000 | 851.7 | 448.8 | 1.897727 | | 9001 | 854.8 | 451.9 | 1.891569 | **AMD EPYC 7352 24-Core Processor** | Size | OldTime | NewTime | Speedup | |------|---------|---------|----------| | 64 | 9.159 | 6.949 | 1.318031 | | 128 | 13.59 | 10.59 | 1.283286 | | 256 | 22.37 | 14.91 | 1.500335 | | 512 | 41.42 | 24.22 | 1.710157 | | 1024 | 81.59 | 45.05 | 1.811099 | | 1500 | 120.4 | 68.35 | 1.761522 | | 2048 | 162.8 | 90.14 | 1.806079 | | 4096 | 321.4 | 180.3 | 1.782585 | | 8192 | 650.4 | 360.8 | 1.802661 | | 9000 | 706.3 | 398.1 | 1.774177 | | 9001 | 712.4 | 398.2 | 1.789051 | Signed-off-by: Tu Dinh Ngoc [Jason: simplified and cleaned up unit tests] Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- tun/checksum.go | 122 +++++++++++++++++++------------------------ tun/checksum_test.go | 63 ++++++++++++++++++++++ 2 files changed, 116 insertions(+), 69 deletions(-) diff --git a/tun/checksum.go b/tun/checksum.go index 29a8fc8fc..b489c56f5 100644 --- a/tun/checksum.go +++ b/tun/checksum.go @@ -1,102 +1,86 @@ package tun -import "encoding/binary" +import ( + "encoding/binary" + "math/bits" +) // TODO: Explore SIMD and/or other assembly optimizations. -// TODO: Test native endian loads. See RFC 1071 section 2 part B. func checksumNoFold(b []byte, initial uint64) uint64 { - ac := initial + tmp := make([]byte, 8) + binary.NativeEndian.PutUint64(tmp, initial) + ac := binary.BigEndian.Uint64(tmp) + var carry uint64 for len(b) >= 128 { - ac += uint64(binary.BigEndian.Uint32(b[:4])) - ac += uint64(binary.BigEndian.Uint32(b[4:8])) - ac += uint64(binary.BigEndian.Uint32(b[8:12])) - ac += uint64(binary.BigEndian.Uint32(b[12:16])) - ac += uint64(binary.BigEndian.Uint32(b[16:20])) - ac += uint64(binary.BigEndian.Uint32(b[20:24])) - ac += uint64(binary.BigEndian.Uint32(b[24:28])) - ac += uint64(binary.BigEndian.Uint32(b[28:32])) - ac += uint64(binary.BigEndian.Uint32(b[32:36])) - ac += uint64(binary.BigEndian.Uint32(b[36:40])) - ac += uint64(binary.BigEndian.Uint32(b[40:44])) - ac += uint64(binary.BigEndian.Uint32(b[44:48])) - ac += uint64(binary.BigEndian.Uint32(b[48:52])) - ac += uint64(binary.BigEndian.Uint32(b[52:56])) - ac += uint64(binary.BigEndian.Uint32(b[56:60])) - ac += uint64(binary.BigEndian.Uint32(b[60:64])) - ac += uint64(binary.BigEndian.Uint32(b[64:68])) - ac += uint64(binary.BigEndian.Uint32(b[68:72])) - ac += uint64(binary.BigEndian.Uint32(b[72:76])) - ac += uint64(binary.BigEndian.Uint32(b[76:80])) - ac += uint64(binary.BigEndian.Uint32(b[80:84])) - ac += uint64(binary.BigEndian.Uint32(b[84:88])) - ac += uint64(binary.BigEndian.Uint32(b[88:92])) - ac += uint64(binary.BigEndian.Uint32(b[92:96])) - ac += uint64(binary.BigEndian.Uint32(b[96:100])) - ac += uint64(binary.BigEndian.Uint32(b[100:104])) - ac += uint64(binary.BigEndian.Uint32(b[104:108])) - ac += uint64(binary.BigEndian.Uint32(b[108:112])) - ac += uint64(binary.BigEndian.Uint32(b[112:116])) - ac += uint64(binary.BigEndian.Uint32(b[116:120])) - ac += uint64(binary.BigEndian.Uint32(b[120:124])) - ac += uint64(binary.BigEndian.Uint32(b[124:128])) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[8:16]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[16:24]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[24:32]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[32:40]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[40:48]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[48:56]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[56:64]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[64:72]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[72:80]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[80:88]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[88:96]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[96:104]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[104:112]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[112:120]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[120:128]), carry) + ac += carry b = b[128:] } if len(b) >= 64 { - ac += uint64(binary.BigEndian.Uint32(b[:4])) - ac += uint64(binary.BigEndian.Uint32(b[4:8])) - ac += uint64(binary.BigEndian.Uint32(b[8:12])) - ac += uint64(binary.BigEndian.Uint32(b[12:16])) - ac += uint64(binary.BigEndian.Uint32(b[16:20])) - ac += uint64(binary.BigEndian.Uint32(b[20:24])) - ac += uint64(binary.BigEndian.Uint32(b[24:28])) - ac += uint64(binary.BigEndian.Uint32(b[28:32])) - ac += uint64(binary.BigEndian.Uint32(b[32:36])) - ac += uint64(binary.BigEndian.Uint32(b[36:40])) - ac += uint64(binary.BigEndian.Uint32(b[40:44])) - ac += uint64(binary.BigEndian.Uint32(b[44:48])) - ac += uint64(binary.BigEndian.Uint32(b[48:52])) - ac += uint64(binary.BigEndian.Uint32(b[52:56])) - ac += uint64(binary.BigEndian.Uint32(b[56:60])) - ac += uint64(binary.BigEndian.Uint32(b[60:64])) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[8:16]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[16:24]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[24:32]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[32:40]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[40:48]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[48:56]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[56:64]), carry) + ac += carry b = b[64:] } if len(b) >= 32 { - ac += uint64(binary.BigEndian.Uint32(b[:4])) - ac += uint64(binary.BigEndian.Uint32(b[4:8])) - ac += uint64(binary.BigEndian.Uint32(b[8:12])) - ac += uint64(binary.BigEndian.Uint32(b[12:16])) - ac += uint64(binary.BigEndian.Uint32(b[16:20])) - ac += uint64(binary.BigEndian.Uint32(b[20:24])) - ac += uint64(binary.BigEndian.Uint32(b[24:28])) - ac += uint64(binary.BigEndian.Uint32(b[28:32])) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[8:16]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[16:24]), carry) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[24:32]), carry) + ac += carry b = b[32:] } if len(b) >= 16 { - ac += uint64(binary.BigEndian.Uint32(b[:4])) - ac += uint64(binary.BigEndian.Uint32(b[4:8])) - ac += uint64(binary.BigEndian.Uint32(b[8:12])) - ac += uint64(binary.BigEndian.Uint32(b[12:16])) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[8:16]), carry) + ac += carry b = b[16:] } if len(b) >= 8 { - ac += uint64(binary.BigEndian.Uint32(b[:4])) - ac += uint64(binary.BigEndian.Uint32(b[4:8])) + ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0) + ac += carry b = b[8:] } if len(b) >= 4 { - ac += uint64(binary.BigEndian.Uint32(b)) + ac, carry = bits.Add64(ac, uint64(binary.NativeEndian.Uint32(b[:4])), 0) + ac += carry b = b[4:] } if len(b) >= 2 { - ac += uint64(binary.BigEndian.Uint16(b)) + ac, carry = bits.Add64(ac, uint64(binary.NativeEndian.Uint16(b[:2])), 0) + ac += carry b = b[2:] } if len(b) == 1 { - ac += uint64(b[0]) << 8 + tmp := binary.NativeEndian.Uint16([]byte{b[0], 0}) + ac, carry = bits.Add64(ac, uint64(tmp), 0) + ac += carry } - return ac + binary.NativeEndian.PutUint64(tmp, ac) + return binary.BigEndian.Uint64(tmp) } func checksum(b []byte, initial uint64) uint16 { diff --git a/tun/checksum_test.go b/tun/checksum_test.go index c1ccff531..4ea9b8b52 100644 --- a/tun/checksum_test.go +++ b/tun/checksum_test.go @@ -1,11 +1,74 @@ package tun import ( + "encoding/binary" "fmt" "math/rand" "testing" + + "golang.org/x/sys/unix" ) +func checksumRef(b []byte, initial uint16) uint16 { + ac := uint64(initial) + + for len(b) >= 2 { + ac += uint64(binary.BigEndian.Uint16(b)) + b = b[2:] + } + if len(b) == 1 { + ac += uint64(b[0]) << 8 + } + + for (ac >> 16) > 0 { + ac = (ac >> 16) + (ac & 0xffff) + } + return uint16(ac) +} + +func pseudoHeaderChecksumRefNoFold(protocol uint8, srcAddr, dstAddr []byte, totalLen uint16) uint16 { + sum := checksumRef(srcAddr, 0) + sum = checksumRef(dstAddr, sum) + sum = checksumRef([]byte{0, protocol}, sum) + tmp := make([]byte, 2) + binary.BigEndian.PutUint16(tmp, totalLen) + return checksumRef(tmp, sum) +} + +func TestChecksum(t *testing.T) { + for length := 0; length <= 9001; length++ { + buf := make([]byte, length) + rng := rand.New(rand.NewSource(1)) + rng.Read(buf) + csum := checksum(buf, 0x1234) + csumRef := checksumRef(buf, 0x1234) + if csum != csumRef { + t.Error("Expected checksum", csumRef, "got", csum) + } + } +} + +func TestPseudoHeaderChecksum(t *testing.T) { + for _, addrLen := range []int{4, 16} { + for length := 0; length <= 9001; length++ { + srcAddr := make([]byte, addrLen) + dstAddr := make([]byte, addrLen) + buf := make([]byte, length) + rng := rand.New(rand.NewSource(1)) + rng.Read(srcAddr) + rng.Read(dstAddr) + rng.Read(buf) + phSum := pseudoHeaderChecksumNoFold(unix.IPPROTO_TCP, srcAddr, dstAddr, uint16(length)) + csum := checksum(buf, phSum) + phSumRef := pseudoHeaderChecksumRefNoFold(unix.IPPROTO_TCP, srcAddr, dstAddr, uint16(length)) + csumRef := checksumRef(buf, phSumRef) + if csum != csumRef { + t.Error("Expected checksumRef", csumRef, "got", csum) + } + } + } +} + func BenchmarkChecksum(b *testing.B) { lengths := []int{ 64, From a193cf4d7c1f7886e0276542669075a01689041a Mon Sep 17 00:00:00 2001 From: ruokeqx Date: Thu, 2 Jan 2025 20:28:33 +0800 Subject: [PATCH 02/34] tun: darwin: fetch flags and mtu from if_msghdr directly Signed-off-by: ruokeqx Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- tun/tun_darwin.go | 34 +++++++++------------------------- 1 file changed, 9 insertions(+), 25 deletions(-) diff --git a/tun/tun_darwin.go b/tun/tun_darwin.go index 407b6f2e9..341afe3c5 100644 --- a/tun/tun_darwin.go +++ b/tun/tun_darwin.go @@ -6,14 +6,12 @@ package tun import ( - "errors" "fmt" "io" "net" "os" "sync" "syscall" - "time" "unsafe" "golang.org/x/sys/unix" @@ -30,18 +28,6 @@ type NativeTun struct { closeOnce sync.Once } -func retryInterfaceByIndex(index int) (iface *net.Interface, err error) { - for i := 0; i < 20; i++ { - iface, err = net.InterfaceByIndex(index) - if err != nil && errors.Is(err, unix.ENOMEM) { - time.Sleep(time.Duration(i) * time.Second / 3) - continue - } - return iface, err - } - return nil, err -} - func (tun *NativeTun) routineRouteListener(tunIfindex int) { var ( statusUp bool @@ -62,26 +48,22 @@ func (tun *NativeTun) routineRouteListener(tunIfindex int) { return } - if n < 14 { + if n < 28 { continue } - if data[3 /* type */] != unix.RTM_IFINFO { + if data[3 /* ifm_type */] != unix.RTM_IFINFO { continue } - ifindex := int(*(*uint16)(unsafe.Pointer(&data[12 /* ifindex */]))) + ifindex := int(*(*uint16)(unsafe.Pointer(&data[12 /* ifm_index */]))) if ifindex != tunIfindex { continue } - iface, err := retryInterfaceByIndex(ifindex) - if err != nil { - tun.errors <- err - return - } + flags := int(*(*uint32)(unsafe.Pointer(&data[8 /* ifm_flags */]))) // Up / Down event - up := (iface.Flags & net.FlagUp) != 0 + up := (flags & syscall.IFF_UP) != 0 if up != statusUp && up { tun.events <- EventUp } @@ -90,11 +72,13 @@ func (tun *NativeTun) routineRouteListener(tunIfindex int) { } statusUp = up + mtu := int(*(*uint32)(unsafe.Pointer(&data[24 /* ifm_data.ifi_mtu */]))) + // MTU changes - if iface.MTU != statusMTU { + if mtu != statusMTU { tun.events <- EventMTUUpdate } - statusMTU = iface.MTU + statusMTU = mtu } } From 7784c5a3622931f920bfcdacaa09bf856b02268f Mon Sep 17 00:00:00 2001 From: Tom Holford Date: Sun, 4 May 2025 18:49:03 +0200 Subject: [PATCH 03/34] global: replaced unused function params with _ Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- conn/errors_default.go | 2 +- conn/features_default.go | 2 +- device/allowedips_test.go | 2 +- device/sticky_default.go | 2 +- device/sticky_linux.go | 4 ++-- 5 files changed, 6 insertions(+), 6 deletions(-) diff --git a/conn/errors_default.go b/conn/errors_default.go index d9675188b..3c9b22357 100644 --- a/conn/errors_default.go +++ b/conn/errors_default.go @@ -7,6 +7,6 @@ package conn -func errShouldDisableUDPGSO(err error) bool { +func errShouldDisableUDPGSO(_ error) bool { return false } diff --git a/conn/features_default.go b/conn/features_default.go index cae2bea52..9fc5088e2 100644 --- a/conn/features_default.go +++ b/conn/features_default.go @@ -10,6 +10,6 @@ package conn import "net" -func supportsUDPOffload(conn *net.UDPConn) (txOffload, rxOffload bool) { +func supportsUDPOffload(_ *net.UDPConn) (txOffload, rxOffload bool) { return } diff --git a/device/allowedips_test.go b/device/allowedips_test.go index 9ef8a761c..0ce45af99 100644 --- a/device/allowedips_test.go +++ b/device/allowedips_test.go @@ -39,7 +39,7 @@ func TestCommonBits(t *testing.T) { } } -func benchmarkTrie(peerNumber, addressNumber, addressLength int, b *testing.B) { +func benchmarkTrie(peerNumber, addressNumber, _ int, b *testing.B) { var trie *trieEntry var peers []*Peer root := parentIndirection{&trie, 2} diff --git a/device/sticky_default.go b/device/sticky_default.go index 10382565e..22e1e15b5 100644 --- a/device/sticky_default.go +++ b/device/sticky_default.go @@ -7,6 +7,6 @@ import ( "golang.zx2c4.com/wireguard/rwcancel" ) -func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, error) { +func (device *Device) startRouteListener(_ conn.Bind) (*rwcancel.RWCancel, error) { return nil, nil } diff --git a/device/sticky_linux.go b/device/sticky_linux.go index 7307b7edb..f23ff0221 100644 --- a/device/sticky_linux.go +++ b/device/sticky_linux.go @@ -9,7 +9,7 @@ * * Currently there is no way to achieve this within the net package: * See e.g. https://github.com/golang/go/issues/17930 - * So this code is remains platform dependent. + * So this code remains platform dependent. */ package device @@ -47,7 +47,7 @@ func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, er return netlinkCancel, nil } -func (device *Device) routineRouteListener(bind conn.Bind, netlinkSock int, netlinkCancel *rwcancel.RWCancel) { +func (device *Device) routineRouteListener(_ conn.Bind, netlinkSock int, netlinkCancel *rwcancel.RWCancel) { type peerEndpointPtr struct { peer *Peer endpoint *conn.Endpoint From 8916471d579c4694323d27513f702a96b899b631 Mon Sep 17 00:00:00 2001 From: Tom Holford Date: Sun, 4 May 2025 18:49:49 +0200 Subject: [PATCH 04/34] device: use rand.NewSource instead of rand.Seed Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- device/allowedips_rand_test.go | 10 +++++----- device/allowedips_test.go | 10 +++++----- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/device/allowedips_rand_test.go b/device/allowedips_rand_test.go index 8dd9b67dd..b863696fb 100644 --- a/device/allowedips_rand_test.go +++ b/device/allowedips_rand_test.go @@ -83,7 +83,7 @@ func TestTrieRandom(t *testing.T) { var peers []*Peer var allowedIPs AllowedIPs - rand.Seed(1) + rng := rand.New(rand.NewSource(1)) for n := 0; n < NumberOfPeers; n++ { peers = append(peers, &Peer{}) @@ -91,14 +91,14 @@ func TestTrieRandom(t *testing.T) { for n := 0; n < NumberOfAddresses; n++ { var addr4 [4]byte - rand.Read(addr4[:]) + rng.Read(addr4[:]) cidr := uint8(rand.Intn(32) + 1) index := rand.Intn(NumberOfPeers) allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom4(addr4), int(cidr)), peers[index]) slow4 = slow4.Insert(addr4[:], cidr, peers[index]) var addr6 [16]byte - rand.Read(addr6[:]) + rng.Read(addr6[:]) cidr = uint8(rand.Intn(128) + 1) index = rand.Intn(NumberOfPeers) allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom16(addr6), int(cidr)), peers[index]) @@ -109,7 +109,7 @@ func TestTrieRandom(t *testing.T) { for p = 0; ; p++ { for n := 0; n < NumberOfTests; n++ { var addr4 [4]byte - rand.Read(addr4[:]) + rng.Read(addr4[:]) peer1 := slow4.Lookup(addr4[:]) peer2 := allowedIPs.Lookup(addr4[:]) if peer1 != peer2 { @@ -117,7 +117,7 @@ func TestTrieRandom(t *testing.T) { } var addr6 [16]byte - rand.Read(addr6[:]) + rng.Read(addr6[:]) peer1 = slow6.Lookup(addr6[:]) peer2 = allowedIPs.Lookup(addr6[:]) if peer1 != peer2 { diff --git a/device/allowedips_test.go b/device/allowedips_test.go index 0ce45af99..7df7da5b8 100644 --- a/device/allowedips_test.go +++ b/device/allowedips_test.go @@ -44,7 +44,7 @@ func benchmarkTrie(peerNumber, addressNumber, _ int, b *testing.B) { var peers []*Peer root := parentIndirection{&trie, 2} - rand.Seed(1) + rng := rand.New(rand.NewSource(1)) const AddressLength = 4 @@ -54,15 +54,15 @@ func benchmarkTrie(peerNumber, addressNumber, _ int, b *testing.B) { for n := 0; n < addressNumber; n++ { var addr [AddressLength]byte - rand.Read(addr[:]) - cidr := uint8(rand.Uint32() % (AddressLength * 8)) - index := rand.Int() % peerNumber + rng.Read(addr[:]) + cidr := uint8(rng.Uint32() % (AddressLength * 8)) + index := rng.Int() % peerNumber root.insert(addr[:], cidr, peers[index]) } for n := 0; n < b.N; n++ { var addr [AddressLength]byte - rand.Read(addr[:]) + rng.Read(addr[:]) trie.lookup(addr[:]) } } From e5a7e5f267b280ac7fc01973ca7d099bb116a941 Mon Sep 17 00:00:00 2001 From: Kurnia D Win Date: Wed, 7 Jun 2023 12:41:02 +0700 Subject: [PATCH 05/34] rwcancel: fix wrong poll event flag on ReadyWrite It should be POLLIN because closeFd is read-only file. Signed-off-by: Kurnia D Win Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- rwcancel/rwcancel.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rwcancel/rwcancel.go b/rwcancel/rwcancel.go index 793e76443..4372453d9 100644 --- a/rwcancel/rwcancel.go +++ b/rwcancel/rwcancel.go @@ -64,7 +64,7 @@ func (rw *RWCancel) ReadyRead() bool { func (rw *RWCancel) ReadyWrite() bool { closeFd := int32(rw.closingReader.Fd()) - pollFds := []unix.PollFd{{Fd: int32(rw.fd), Events: unix.POLLOUT}, {Fd: closeFd, Events: unix.POLLOUT}} + pollFds := []unix.PollFd{{Fd: int32(rw.fd), Events: unix.POLLOUT}, {Fd: closeFd, Events: unix.POLLIN}} var err error for { _, err = unix.Poll(pollFds, -1) From e49aab52cadf5c6605d1046879de8c549fc9730f Mon Sep 17 00:00:00 2001 From: Alexander Yastrebov Date: Thu, 26 Dec 2024 20:36:53 +0100 Subject: [PATCH 06/34] device: reduce RoutineHandshake allocations MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Reduce allocations by eliminating byte reader, hand-rolled decoding and reusing message structs. Synthetic benchmark: var msgSink MessageInitiation func BenchmarkMessageInitiationUnmarshal(b *testing.B) { packet := make([]byte, MessageInitiationSize) reader := bytes.NewReader(packet) err := binary.Read(reader, binary.LittleEndian, &msgSink) if err != nil { b.Fatal(err) } b.Run("binary.Read", func(b *testing.B) { b.ReportAllocs() for range b.N { reader := bytes.NewReader(packet) _ = binary.Read(reader, binary.LittleEndian, &msgSink) } }) b.Run("unmarshal", func(b *testing.B) { b.ReportAllocs() for range b.N { _ = msgSink.unmarshal(packet) } }) } Results: │ - │ │ sec/op │ MessageInitiationUnmarshal/binary.Read-8 1.508µ ± 2% MessageInitiationUnmarshal/unmarshal-8 12.66n ± 2% │ - │ │ B/op │ MessageInitiationUnmarshal/binary.Read-8 208.0 ± 0% MessageInitiationUnmarshal/unmarshal-8 0.000 ± 0% │ - │ │ allocs/op │ MessageInitiationUnmarshal/binary.Read-8 2.000 ± 0% MessageInitiationUnmarshal/unmarshal-8 0.000 ± 0% Signed-off-by: Alexander Yastrebov Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- device/noise-protocol.go | 48 ++++++++++++++++++++++++++++++++++++++++ device/receive.go | 10 +++------ 2 files changed, 51 insertions(+), 7 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index b72acf85d..12368ec62 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -6,6 +6,7 @@ package device import ( + "encoding/binary" "errors" "fmt" "sync" @@ -115,6 +116,53 @@ type MessageCookieReply struct { Cookie [blake2s.Size128 + poly1305.TagSize]byte } +var errMessageTooShort = errors.New("message too short") + +func (msg *MessageInitiation) unmarshal(b []byte) error { + if len(b) < MessageInitiationSize { + return errMessageTooShort + } + + msg.Type = binary.LittleEndian.Uint32(b) + msg.Sender = binary.LittleEndian.Uint32(b[4:]) + copy(msg.Ephemeral[:], b[8:]) + copy(msg.Static[:], b[8+len(msg.Ephemeral):]) + copy(msg.Timestamp[:], b[8+len(msg.Ephemeral)+len(msg.Static):]) + copy(msg.MAC1[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.Timestamp):]) + copy(msg.MAC2[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.Timestamp)+len(msg.MAC1):]) + + return nil +} + +func (msg *MessageResponse) unmarshal(b []byte) error { + if len(b) < MessageResponseSize { + return errMessageTooShort + } + + msg.Type = binary.LittleEndian.Uint32(b) + msg.Sender = binary.LittleEndian.Uint32(b[4:]) + msg.Receiver = binary.LittleEndian.Uint32(b[8:]) + copy(msg.Ephemeral[:], b[12:]) + copy(msg.Empty[:], b[12+len(msg.Ephemeral):]) + copy(msg.MAC1[:], b[12+len(msg.Ephemeral)+len(msg.Empty):]) + copy(msg.MAC2[:], b[12+len(msg.Ephemeral)+len(msg.Empty)+len(msg.MAC1):]) + + return nil +} + +func (msg *MessageCookieReply) unmarshal(b []byte) error { + if len(b) < MessageCookieReplySize { + return errMessageTooShort + } + + msg.Type = binary.LittleEndian.Uint32(b) + msg.Receiver = binary.LittleEndian.Uint32(b[4:]) + copy(msg.Nonce[:], b[8:]) + copy(msg.Cookie[:], b[8+len(msg.Nonce):]) + + return nil +} + type Handshake struct { state handshakeState mutex sync.RWMutex diff --git a/device/receive.go b/device/receive.go index c7b6c87fc..13929577e 100644 --- a/device/receive.go +++ b/device/receive.go @@ -6,7 +6,6 @@ package device import ( - "bytes" "encoding/binary" "errors" "net" @@ -287,8 +286,7 @@ func (device *Device) RoutineHandshake(id int) { // unmarshal packet var reply MessageCookieReply - reader := bytes.NewReader(elem.packet) - err := binary.Read(reader, binary.LittleEndian, &reply) + err := reply.unmarshal(elem.packet) if err != nil { device.log.Verbosef("Failed to decode cookie reply") goto skip @@ -353,8 +351,7 @@ func (device *Device) RoutineHandshake(id int) { // unmarshal var msg MessageInitiation - reader := bytes.NewReader(elem.packet) - err := binary.Read(reader, binary.LittleEndian, &msg) + err := msg.unmarshal(elem.packet) if err != nil { device.log.Errorf("Failed to decode initiation message") goto skip @@ -386,8 +383,7 @@ func (device *Device) RoutineHandshake(id int) { // unmarshal var msg MessageResponse - reader := bytes.NewReader(elem.packet) - err := binary.Read(reader, binary.LittleEndian, &msg) + err := msg.unmarshal(elem.packet) if err != nil { device.log.Errorf("Failed to decode response message") goto skip From 8f357d81cff2675eaf3d43ef5014e1da563d6bcc Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Thu, 15 May 2025 16:48:14 +0200 Subject: [PATCH 07/34] device: make unmarshall length checks exact This is already enforced in receive.go, but if these unmarshallers are to have error return values anyway, make them as explicit as possible. Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- device/noise-protocol.go | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 12368ec62..5f713ee5e 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -116,11 +116,11 @@ type MessageCookieReply struct { Cookie [blake2s.Size128 + poly1305.TagSize]byte } -var errMessageTooShort = errors.New("message too short") +var errMessageLengthMismatch = errors.New("message length mismatch") func (msg *MessageInitiation) unmarshal(b []byte) error { - if len(b) < MessageInitiationSize { - return errMessageTooShort + if len(b) != MessageInitiationSize { + return errMessageLengthMismatch } msg.Type = binary.LittleEndian.Uint32(b) @@ -135,8 +135,8 @@ func (msg *MessageInitiation) unmarshal(b []byte) error { } func (msg *MessageResponse) unmarshal(b []byte) error { - if len(b) < MessageResponseSize { - return errMessageTooShort + if len(b) != MessageResponseSize { + return errMessageLengthMismatch } msg.Type = binary.LittleEndian.Uint32(b) @@ -151,8 +151,8 @@ func (msg *MessageResponse) unmarshal(b []byte) error { } func (msg *MessageCookieReply) unmarshal(b []byte) error { - if len(b) < MessageCookieReplySize { - return errMessageTooShort + if len(b) != MessageCookieReplySize { + return errMessageLengthMismatch } msg.Type = binary.LittleEndian.Uint32(b) From ae5a1c66ce7af991e71bbcc697c8d5c1ade923ad Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Thu, 15 May 2025 16:54:03 +0200 Subject: [PATCH 08/34] version: bump snapshot Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- version.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/version.go b/version.go index db75bb938..80f2d4ba7 100644 --- a/version.go +++ b/version.go @@ -1,3 +1,3 @@ package main -const Version = "0.0.20230223" +const Version = "0.0.20250515" From 87b78d5a5a6b908ed8afa3f3e2331ccb62c1f390 Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Tue, 20 May 2025 23:03:06 +0200 Subject: [PATCH 09/34] device: add support for removing allowedips individually This pairs with the recent change in wireguard-tools. Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- device/allowedips.go | 87 +++++++++++++++++++++++++-------------- device/allowedips_test.go | 57 +++++++++++++++++++++++++ device/uapi.go | 15 ++++++- 3 files changed, 125 insertions(+), 34 deletions(-) diff --git a/device/allowedips.go b/device/allowedips.go index b40c8170c..d15373cfe 100644 --- a/device/allowedips.go +++ b/device/allowedips.go @@ -223,6 +223,60 @@ func (table *AllowedIPs) EntriesForPeer(peer *Peer, cb func(prefix netip.Prefix) } } +func (node *trieEntry) remove() { + node.removeFromPeerEntries() + node.peer = nil + if node.child[0] != nil && node.child[1] != nil { + return + } + bit := 0 + if node.child[0] == nil { + bit = 1 + } + child := node.child[bit] + if child != nil { + child.parent = node.parent + } + *node.parent.parentBit = child + if node.child[0] != nil || node.child[1] != nil || node.parent.parentBitType > 1 { + node.zeroizePointers() + return + } + parent := (*trieEntry)(unsafe.Pointer(uintptr(unsafe.Pointer(node.parent.parentBit)) - unsafe.Offsetof(node.child) - unsafe.Sizeof(node.child[0])*uintptr(node.parent.parentBitType))) + if parent.peer != nil { + node.zeroizePointers() + return + } + child = parent.child[node.parent.parentBitType^1] + if child != nil { + child.parent = parent.parent + } + *parent.parent.parentBit = child + node.zeroizePointers() + parent.zeroizePointers() +} + +func (table *AllowedIPs) Remove(prefix netip.Prefix, peer *Peer) { + table.mutex.Lock() + defer table.mutex.Unlock() + var node *trieEntry + var exact bool + + if prefix.Addr().Is6() { + ip := prefix.Addr().As16() + node, exact = table.IPv6.nodePlacement(ip[:], uint8(prefix.Bits())) + } else if prefix.Addr().Is4() { + ip := prefix.Addr().As4() + node, exact = table.IPv4.nodePlacement(ip[:], uint8(prefix.Bits())) + } else { + panic(errors.New("removing unknown address type")) + } + if !exact || node == nil || peer != node.peer { + return + } + node.remove() +} + func (table *AllowedIPs) RemoveByPeer(peer *Peer) { table.mutex.Lock() defer table.mutex.Unlock() @@ -230,38 +284,7 @@ func (table *AllowedIPs) RemoveByPeer(peer *Peer) { var next *list.Element for elem := peer.trieEntries.Front(); elem != nil; elem = next { next = elem.Next() - node := elem.Value.(*trieEntry) - - node.removeFromPeerEntries() - node.peer = nil - if node.child[0] != nil && node.child[1] != nil { - continue - } - bit := 0 - if node.child[0] == nil { - bit = 1 - } - child := node.child[bit] - if child != nil { - child.parent = node.parent - } - *node.parent.parentBit = child - if node.child[0] != nil || node.child[1] != nil || node.parent.parentBitType > 1 { - node.zeroizePointers() - continue - } - parent := (*trieEntry)(unsafe.Pointer(uintptr(unsafe.Pointer(node.parent.parentBit)) - unsafe.Offsetof(node.child) - unsafe.Sizeof(node.child[0])*uintptr(node.parent.parentBitType))) - if parent.peer != nil { - node.zeroizePointers() - continue - } - child = parent.child[node.parent.parentBitType^1] - if child != nil { - child.parent = parent.parent - } - *parent.parent.parentBit = child - node.zeroizePointers() - parent.zeroizePointers() + elem.Value.(*trieEntry).remove() } } diff --git a/device/allowedips_test.go b/device/allowedips_test.go index 7df7da5b8..a4b08a399 100644 --- a/device/allowedips_test.go +++ b/device/allowedips_test.go @@ -101,6 +101,10 @@ func TestTrieIPv4(t *testing.T) { allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom4([4]byte{a, b, c, d}), int(cidr)), peer) } + remove := func(peer *Peer, a, b, c, d byte, cidr uint8) { + allowedIPs.Remove(netip.PrefixFrom(netip.AddrFrom4([4]byte{a, b, c, d}), int(cidr)), peer) + } + assertEQ := func(peer *Peer, a, b, c, d byte) { p := allowedIPs.Lookup([]byte{a, b, c, d}) if p != peer { @@ -176,6 +180,21 @@ func TestTrieIPv4(t *testing.T) { allowedIPs.RemoveByPeer(a) assertNEQ(a, 192, 168, 0, 1) + + insert(a, 1, 0, 0, 0, 32) + insert(a, 192, 0, 0, 0, 24) + assertEQ(a, 1, 0, 0, 0) + assertEQ(a, 192, 0, 0, 1) + remove(a, 192, 0, 0, 0, 32) + assertEQ(a, 192, 0, 0, 1) + remove(nil, 192, 0, 0, 0, 24) + assertEQ(a, 192, 0, 0, 1) + remove(b, 192, 0, 0, 0, 24) + assertEQ(a, 192, 0, 0, 1) + remove(a, 192, 0, 0, 0, 24) + assertNEQ(a, 192, 0, 0, 1) + remove(a, 1, 0, 0, 0, 32) + assertNEQ(a, 1, 0, 0, 0) } /* Test ported from kernel implementation: @@ -211,6 +230,15 @@ func TestTrieIPv6(t *testing.T) { allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom16(*(*[16]byte)(addr)), int(cidr)), peer) } + remove := func(peer *Peer, a, b, c, d uint32, cidr uint8) { + var addr []byte + addr = append(addr, expand(a)...) + addr = append(addr, expand(b)...) + addr = append(addr, expand(c)...) + addr = append(addr, expand(d)...) + allowedIPs.Remove(netip.PrefixFrom(netip.AddrFrom16(*(*[16]byte)(addr)), int(cidr)), peer) + } + assertEQ := func(peer *Peer, a, b, c, d uint32) { var addr []byte addr = append(addr, expand(a)...) @@ -223,6 +251,18 @@ func TestTrieIPv6(t *testing.T) { } } + assertNEQ := func(peer *Peer, a, b, c, d uint32) { + var addr []byte + addr = append(addr, expand(a)...) + addr = append(addr, expand(b)...) + addr = append(addr, expand(c)...) + addr = append(addr, expand(d)...) + p := allowedIPs.Lookup(addr) + if p == peer { + t.Error("Assert NEQ failed") + } + } + insert(d, 0x26075300, 0x60006b00, 0, 0xc05f0543, 128) insert(c, 0x26075300, 0x60006b00, 0, 0, 64) insert(e, 0, 0, 0, 0, 0) @@ -244,4 +284,21 @@ func TestTrieIPv6(t *testing.T) { assertEQ(h, 0x24046800, 0x40040800, 0, 0) assertEQ(h, 0x24046800, 0x40040800, 0x10101010, 0x10101010) assertEQ(a, 0x24046800, 0x40040800, 0xdeadbeef, 0xdeadbeef) + + insert(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128) + insert(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0, 98) + assertEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef) + assertEQ(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0x10101010) + remove(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 96) + assertEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef) + remove(nil, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128) + assertEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef) + remove(b, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128) + assertEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef) + remove(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128) + assertNEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef) + remove(b, 0x24446800, 0xf0e40800, 0xeeaebeef, 0, 98) + assertEQ(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0x10101010) + remove(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0, 98) + assertNEQ(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0x10101010) } diff --git a/device/uapi.go b/device/uapi.go index 521a7411c..cc69488b4 100644 --- a/device/uapi.go +++ b/device/uapi.go @@ -371,7 +371,14 @@ func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error device.allowedips.RemoveByPeer(peer.Peer) case "allowed_ip": - device.log.Verbosef("%v - UAPI: Adding allowedip", peer.Peer) + add := true + verb := "Adding" + if len(value) > 0 && value[0] == '-' { + add = false + verb = "Removing" + value = value[1:] + } + device.log.Verbosef("%v - UAPI: %s allowedip", peer.Peer, verb) prefix, err := netip.ParsePrefix(value) if err != nil { return ipcErrorf(ipc.IpcErrorInvalid, "failed to set allowed ip: %w", err) @@ -379,7 +386,11 @@ func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error if peer.dummy { return nil } - device.allowedips.Insert(prefix, peer.Peer) + if add { + device.allowedips.Insert(prefix, peer.Peer) + } else { + device.allowedips.Remove(prefix, peer.Peer) + } case "protocol_version": if value != "1" { From 3d12258a22fe6464a9dd7bdb0944517748885bb7 Mon Sep 17 00:00:00 2001 From: Alexander Yastrebov Date: Sat, 17 May 2025 11:34:30 +0200 Subject: [PATCH 10/34] device: optimize message encoding MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Optimize message encoding by eliminating binary.Write (which internally uses reflection) in favour of hand-rolled encoding. This is companion to 9e7529c3d2d0c54f4d5384c01645a9279e4740ae. Synthetic benchmark: var packetSink []byte func BenchmarkMessageInitiationMarshal(b *testing.B) { var msg MessageInitiation b.Run("binary.Write", func(b *testing.B) { b.ReportAllocs() for range b.N { var buf [MessageInitiationSize]byte writer := bytes.NewBuffer(buf[:0]) _ = binary.Write(writer, binary.LittleEndian, msg) packetSink = writer.Bytes() } }) b.Run("binary.Encode", func(b *testing.B) { b.ReportAllocs() for range b.N { packet := make([]byte, MessageInitiationSize) _, _ = binary.Encode(packet, binary.LittleEndian, msg) packetSink = packet } }) b.Run("marshal", func(b *testing.B) { b.ReportAllocs() for range b.N { packet := make([]byte, MessageInitiationSize) _ = msg.marshal(packet) packetSink = packet } }) } Results: │ - │ │ sec/op │ MessageInitiationMarshal/binary.Write-8 1.337µ ± 0% MessageInitiationMarshal/binary.Encode-8 1.242µ ± 0% MessageInitiationMarshal/marshal-8 53.05n ± 1% │ - │ │ B/op │ MessageInitiationMarshal/binary.Write-8 368.0 ± 0% MessageInitiationMarshal/binary.Encode-8 160.0 ± 0% MessageInitiationMarshal/marshal-8 160.0 ± 0% │ - │ │ allocs/op │ MessageInitiationMarshal/binary.Write-8 3.000 ± 0% MessageInitiationMarshal/binary.Encode-8 1.000 ± 0% MessageInitiationMarshal/marshal-8 1.000 ± 0% Signed-off-by: Alexander Yastrebov Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- device/noise-protocol.go | 45 ++++++++++++++++++++++++++++++++++++++++ device/send.go | 21 +++++++------------ 2 files changed, 53 insertions(+), 13 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 5f713ee5e..5cf1702b6 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -134,6 +134,22 @@ func (msg *MessageInitiation) unmarshal(b []byte) error { return nil } +func (msg *MessageInitiation) marshal(b []byte) error { + if len(b) != MessageInitiationSize { + return errMessageLengthMismatch + } + + binary.LittleEndian.PutUint32(b, msg.Type) + binary.LittleEndian.PutUint32(b[4:], msg.Sender) + copy(b[8:], msg.Ephemeral[:]) + copy(b[8+len(msg.Ephemeral):], msg.Static[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static):], msg.Timestamp[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.Timestamp):], msg.MAC1[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.Timestamp)+len(msg.MAC1):], msg.MAC2[:]) + + return nil +} + func (msg *MessageResponse) unmarshal(b []byte) error { if len(b) != MessageResponseSize { return errMessageLengthMismatch @@ -150,6 +166,22 @@ func (msg *MessageResponse) unmarshal(b []byte) error { return nil } +func (msg *MessageResponse) marshal(b []byte) error { + if len(b) != MessageResponseSize { + return errMessageLengthMismatch + } + + binary.LittleEndian.PutUint32(b, msg.Type) + binary.LittleEndian.PutUint32(b[4:], msg.Sender) + binary.LittleEndian.PutUint32(b[8:], msg.Receiver) + copy(b[12:], msg.Ephemeral[:]) + copy(b[12+len(msg.Ephemeral):], msg.Empty[:]) + copy(b[12+len(msg.Ephemeral)+len(msg.Empty):], msg.MAC1[:]) + copy(b[12+len(msg.Ephemeral)+len(msg.Empty)+len(msg.MAC1):], msg.MAC2[:]) + + return nil +} + func (msg *MessageCookieReply) unmarshal(b []byte) error { if len(b) != MessageCookieReplySize { return errMessageLengthMismatch @@ -163,6 +195,19 @@ func (msg *MessageCookieReply) unmarshal(b []byte) error { return nil } +func (msg *MessageCookieReply) marshal(b []byte) error { + if len(b) != MessageCookieReplySize { + return errMessageLengthMismatch + } + + binary.LittleEndian.PutUint32(b, msg.Type) + binary.LittleEndian.PutUint32(b[4:], msg.Receiver) + copy(b[8:], msg.Nonce[:]) + copy(b[8+len(msg.Nonce):], msg.Cookie[:]) + + return nil +} + type Handshake struct { state handshakeState mutex sync.RWMutex diff --git a/device/send.go b/device/send.go index 38f55c2a4..ff8f7da50 100644 --- a/device/send.go +++ b/device/send.go @@ -6,7 +6,6 @@ package device import ( - "bytes" "encoding/binary" "errors" "net" @@ -124,10 +123,8 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error { return err } - var buf [MessageInitiationSize]byte - writer := bytes.NewBuffer(buf[:0]) - binary.Write(writer, binary.LittleEndian, msg) - packet := writer.Bytes() + packet := make([]byte, MessageInitiationSize) + _ = msg.marshal(packet) peer.cookieGenerator.AddMacs(packet) peer.timersAnyAuthenticatedPacketTraversal() @@ -155,10 +152,8 @@ func (peer *Peer) SendHandshakeResponse() error { return err } - var buf [MessageResponseSize]byte - writer := bytes.NewBuffer(buf[:0]) - binary.Write(writer, binary.LittleEndian, response) - packet := writer.Bytes() + packet := make([]byte, MessageResponseSize) + _ = response.marshal(packet) peer.cookieGenerator.AddMacs(packet) err = peer.BeginSymmetricSession() @@ -189,11 +184,11 @@ func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement) return err } - var buf [MessageCookieReplySize]byte - writer := bytes.NewBuffer(buf[:0]) - binary.Write(writer, binary.LittleEndian, reply) + packet := make([]byte, MessageCookieReplySize) + _ = reply.marshal(packet) // TODO: allocation could be avoided - device.net.bind.Send([][]byte{writer.Bytes()}, initiatingElem.endpoint) + device.net.bind.Send([][]byte{packet}, initiatingElem.endpoint) + return nil } From 56a350659c189e3ac4f7847d4ff92173c1d660ec Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Thu, 22 May 2025 01:33:55 +0200 Subject: [PATCH 11/34] conn: don't enable GRO on Linux < 5.12 Kernels below 5.12 are missing this: commit 98184612aca0a9ee42b8eb0262a49900ee9eef0d Author: Norman Maurer Date: Thu Apr 1 08:59:17 2021 net: udp: Add support for getsockopt(..., ..., UDP_GRO, ..., ...); Support for UDP_GRO was added in the past but the implementation for getsockopt was missed which did lead to an error when we tried to retrieve the setting for UDP_GRO. This patch adds the missing switch case for UDP_GRO Fixes: e20cf8d3f1f7 ("udp: implement GRO for plain UDP sockets.") Signed-off-by: Norman Maurer Reviewed-by: David Ahern Signed-off-by: David S. Miller That means we can't set the option and then read it back later. Given how buggy UDP_GRO is in general on odd kernels, just disable it on older kernels all together. Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- conn/controlfns_linux.go | 40 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/conn/controlfns_linux.go b/conn/controlfns_linux.go index 3447349f7..f0deefaaf 100644 --- a/conn/controlfns_linux.go +++ b/conn/controlfns_linux.go @@ -13,6 +13,35 @@ import ( "golang.org/x/sys/unix" ) +// Taken from go/src/internal/syscall/unix/kernel_version_linux.go +func kernelVersion() (major, minor int) { + var uname unix.Utsname + if err := unix.Uname(&uname); err != nil { + return + } + + var ( + values [2]int + value, vi int + ) + for _, c := range uname.Release { + if '0' <= c && c <= '9' { + value = (value * 10) + int(c-'0') + } else { + // Note that we're assuming N.N.N here. + // If we see anything else, we are likely to mis-parse it. + values[vi] = value + vi++ + if vi >= len(values) { + break + } + value = 0 + } + } + + return values[0], values[1] +} + func init() { controlFns = append(controlFns, @@ -60,6 +89,17 @@ func init() { // Attempt to enable UDP_GRO func(network, address string, c syscall.RawConn) error { + // Kernels below 5.12 are missing 98184612aca0 ("net: + // udp: Add support for getsockopt(..., ..., UDP_GRO, + // ..., ...);"), which means we can't read this back + // later. We could pipe the return value through to + // the rest of the code, but UDP_GRO is kind of buggy + // anyway, so just gate this here. + major, minor := kernelVersion() + if major < 5 || (major == 5 && minor < 12) { + return nil + } + c.Control(func(fd uintptr) { _ = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO, 1) }) From 853ff703e4d89342ac398f3a1bada575104c1eda Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Thu, 22 May 2025 01:45:02 +0200 Subject: [PATCH 12/34] version: bump snapshot Signed-off-by: Jason A. Donenfeld Signed-off-by: Mateus Franco --- version.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/version.go b/version.go index 80f2d4ba7..d5524e88c 100644 --- a/version.go +++ b/version.go @@ -1,3 +1,3 @@ package main -const Version = "0.0.20250515" +const Version = "0.0.20250522" From 4882a13c5d5820e4d6da455b1e33c7638fefa003 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sat, 30 Aug 2025 14:16:08 -0300 Subject: [PATCH 13/34] chore: add indirect dependency for cloudflare/circl v1.6.1 Signed-off-by: Mateus Franco --- go.mod | 1 + go.sum | 2 ++ 2 files changed, 3 insertions(+) diff --git a/go.mod b/go.mod index 2a80e0001..6947698ca 100644 --- a/go.mod +++ b/go.mod @@ -11,6 +11,7 @@ require ( ) require ( + github.com/cloudflare/circl v1.6.1 // indirect github.com/google/btree v1.1.2 // indirect golang.org/x/time v0.7.0 // indirect ) diff --git a/go.sum b/go.sum index 61875c160..a7bc79b1c 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,5 @@ +github.com/cloudflare/circl v1.6.1 h1:zqIqSPIndyBh1bjLVVDHMPpVKqp8Su/V+6MeDzzQBQ0= +github.com/cloudflare/circl v1.6.1/go.mod h1:uddAzsPgqdMAYatqJ0lsjX1oECcQLIlRpzZh3pJrofs= github.com/google/btree v1.1.2 h1:xf4v41cLI2Z6FxbKm+8Bu+m8ifhj15JuZ9sa0jZCMUU= github.com/google/btree v1.1.2/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4= golang.org/x/crypto v0.37.0 h1:kJNSjF/Xp7kU0iB2Z+9viTPMW4EqqsrywMXLJOOsXSE= From d5069657fbc9855b9fb3a16932e20d99d5bb997c Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sat, 30 Aug 2025 14:51:44 -0300 Subject: [PATCH 14/34] feat: add ML-KEM post-quantum key types Signed-off-by: Mateus Franco --- device/noise-types.go | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/device/noise-types.go b/device/noise-types.go index 41c944e14..523ffa12c 100644 --- a/device/noise-types.go +++ b/device/noise-types.go @@ -76,3 +76,14 @@ func (key NoisePublicKey) Equals(tar NoisePublicKey) bool { func (key *NoisePresharedKey) FromHex(src string) error { return loadExactHex(key[:], src) } + +const ( + MLKEMPublicKeySize = 1568 + MLKEMPrivateKeySize = 3168 + MLKEMCiphertextSize = 1568 +) + +type ( + MLKEMPublicKey [MLKEMPublicKeySize]byte + MLKEMPrivateKey [MLKEMPrivateKeySize]byte +) From 8851ccd208e0ec9d55590955a877d05d470cac53 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sat, 30 Aug 2025 14:53:30 -0300 Subject: [PATCH 15/34] feat: (device) add ML-KEM public key to peer handshake struct Signed-off-by: Mateus Franco --- device/noise-protocol.go | 1 + 1 file changed, 1 insertion(+) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 5cf1702b6..a4c7090fc 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -218,6 +218,7 @@ type Handshake struct { localIndex uint32 // used to clear hash-table remoteIndex uint32 // index for sending remoteStatic NoisePublicKey // long term key + remoteMLKEMStatic MLKEMPublicKey // long term remote ML-KEM static public key remoteEphemeral NoisePublicKey // ephemeral public key precomputedStaticStatic [NoisePublicKeySize]byte // precomputed shared secret lastTimestamp tai64n.Timestamp From e66ed2c1a3567e47d9f93ca85c34b21cf113c1d8 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sat, 30 Aug 2025 14:54:56 -0300 Subject: [PATCH 16/34] feat(device): Add ML-KEM key pair fields to staticIdentity Signed-off-by: Mateus Franco --- device/device.go | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/device/device.go b/device/device.go index 6854ed85a..d889a9b81 100644 --- a/device/device.go +++ b/device/device.go @@ -49,8 +49,10 @@ type Device struct { staticIdentity struct { sync.RWMutex - privateKey NoisePrivateKey - publicKey NoisePublicKey + privateKey NoisePrivateKey + publicKey NoisePublicKey + mlkemPrivateKey MLKEMPrivateKey + mlkemPublicKey MLKEMPublicKey } peers struct { From 9e84b95bc74d895ec6f969e30b5be47f12e83b8d Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sat, 30 Aug 2025 15:14:50 -0300 Subject: [PATCH 17/34] feat(device): Update MessageInitiation to include ML-KEM ciphertext size Signed-off-by: Mateus Franco --- device/noise-protocol.go | 29 ++++++++++++++++------------- 1 file changed, 16 insertions(+), 13 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index a4c7090fc..820d9506f 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -61,13 +61,13 @@ const ( ) const ( - MessageInitiationSize = 148 // size of handshake initiation message - MessageResponseSize = 92 // size of response message - MessageCookieReplySize = 64 // size of cookie reply message - MessageTransportHeaderSize = 16 // size of data preceding content in transport message - MessageTransportSize = MessageTransportHeaderSize + poly1305.TagSize // size of empty transport - MessageKeepaliveSize = MessageTransportSize // size of keepalive - MessageHandshakeSize = MessageInitiationSize // size of largest handshake related message + MessageInitiationSize = 148 + (MLKEMCiphertextSize + poly1305.TagSize) // size of handshake initiation message + MessageResponseSize = 92 // size of response message + MessageCookieReplySize = 64 // size of cookie reply message + MessageTransportHeaderSize = 16 // size of data preceding content in transport message + MessageTransportSize = MessageTransportHeaderSize + poly1305.TagSize // size of empty transport + MessageKeepaliveSize = MessageTransportSize // size of keepalive + MessageHandshakeSize = MessageInitiationSize // size of largest handshake related message ) const ( @@ -87,6 +87,7 @@ type MessageInitiation struct { Sender uint32 Ephemeral NoisePublicKey Static [NoisePublicKeySize + poly1305.TagSize]byte + MLKEM [MLKEMCiphertextSize + poly1305.TagSize]byte Timestamp [tai64n.TimestampSize + poly1305.TagSize]byte MAC1 [blake2s.Size128]byte MAC2 [blake2s.Size128]byte @@ -127,9 +128,10 @@ func (msg *MessageInitiation) unmarshal(b []byte) error { msg.Sender = binary.LittleEndian.Uint32(b[4:]) copy(msg.Ephemeral[:], b[8:]) copy(msg.Static[:], b[8+len(msg.Ephemeral):]) - copy(msg.Timestamp[:], b[8+len(msg.Ephemeral)+len(msg.Static):]) - copy(msg.MAC1[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.Timestamp):]) - copy(msg.MAC2[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.Timestamp)+len(msg.MAC1):]) + copy(msg.MLKEM[:], b[8+len(msg.Ephemeral)+len(msg.Static):]) + copy(msg.Timestamp[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM):]) + copy(msg.MAC1[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp):]) + copy(msg.MAC2[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp)+len(msg.MAC1):]) return nil } @@ -143,9 +145,10 @@ func (msg *MessageInitiation) marshal(b []byte) error { binary.LittleEndian.PutUint32(b[4:], msg.Sender) copy(b[8:], msg.Ephemeral[:]) copy(b[8+len(msg.Ephemeral):], msg.Static[:]) - copy(b[8+len(msg.Ephemeral)+len(msg.Static):], msg.Timestamp[:]) - copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.Timestamp):], msg.MAC1[:]) - copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.Timestamp)+len(msg.MAC1):], msg.MAC2[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static):], msg.MLKEM[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM):], msg.Timestamp[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp):], msg.MAC1[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp)+len(msg.MAC1):], msg.MAC2[:]) return nil } From a04398abfa9bdf6d323d9c3cc4dcaff455399f01 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sun, 31 Aug 2025 11:43:31 -0300 Subject: [PATCH 18/34] feat(device): Add post-quantum encapsulation in CreateMessageInitiation Signed-off-by: Mateus Franco --- device/noise-protocol.go | 32 +++++++++++++++++++++++++++++--- 1 file changed, 29 insertions(+), 3 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 820d9506f..59f0041b2 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -12,6 +12,7 @@ import ( "sync" "time" + "github.com/cloudflare/circl/kem/kyber/kyber1024" "golang.org/x/crypto/blake2s" "golang.org/x/crypto/chacha20poly1305" "golang.org/x/crypto/poly1305" @@ -298,19 +299,44 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e handshake.mixKey(msg.Ephemeral[:]) handshake.mixHash(msg.Ephemeral[:]) + // post-quantum encapsulation + scheme := kyber1024.Scheme() + pk, err := scheme.UnmarshalBinaryPublicKey(handshake.remoteMLKEMStatic[:]) + if err != nil { + return nil, err + } + + ciphertext, mlkemSecret, err := scheme.Encapsulate(pk) + if err != nil { + return nil, err + } + + // encrypy KEM ciphertext and mix into handshake hash + var key [chacha20poly1305.KeySize]byte + KDF1(&key, handshake.chainKey[:], []byte("pqc-ciphertext-key")) + aead, _ := chacha20poly1305.New(key[:]) + aead.Seal(msg.MLKEM[:0], ZeroNonce[:], ciphertext, handshake.hash[:]) + handshake.mixHash(msg.MLKEM[:]) + // encrypt static key ss, err := handshake.localEphemeral.sharedSecret(handshake.remoteStatic) if err != nil { return nil, err } - var key [chacha20poly1305.KeySize]byte + + // mix classic (ss) and post-quantum secret (mlkemSecret) into a single secret + var combinedSecret [blake2s.Size]byte + KDF2(&combinedSecret, nil, ss[:], mlkemSecret) + + // from this point on, use the combined secret to feed the Noise KDF KDF2( &handshake.chainKey, &key, handshake.chainKey[:], - ss[:], + combinedSecret[:], ) - aead, _ := chacha20poly1305.New(key[:]) + + aead, _ = chacha20poly1305.New(key[:]) aead.Seal(msg.Static[:0], ZeroNonce[:], device.staticIdentity.publicKey[:], handshake.hash[:]) handshake.mixHash(msg.Static[:]) From aadc8a4faec4470fe4a1929be00d467dac73e4c8 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sun, 31 Aug 2025 11:48:16 -0300 Subject: [PATCH 19/34] feat(device): Implement post-quantum decapsulation in ConsumeMessageInitiation Signed-off-by: Mateus Franco --- device/noise-protocol.go | 46 +++++++++++++++++++++++++++++++--------- 1 file changed, 36 insertions(+), 10 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 59f0041b2..22c579184 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -391,7 +391,9 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { if err != nil { return nil } - KDF2(&chainKey, &key, chainKey[:], ss[:]) + + var tempChainKey [blake2s.Size]byte + KDF2(&tempChainKey, &key, chainKey[:], ss[:]) aead, _ := chacha20poly1305.New(key[:]) _, err = aead.Open(peerPK[:0], ZeroNonce[:], msg.Static[:], hash[:]) if err != nil { @@ -400,7 +402,6 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { mixHash(&hash, &hash, msg.Static[:]) // lookup peer - peer := device.LookupPeer(peerPK) if peer == nil || !peer.isRunning.Load() { return nil @@ -408,12 +409,39 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { handshake := &peer.handshake - // verify identity + // decrypt KEM ciphertext + KDF1(&key, chainKey[:], []byte("pqc-ciphertext-key")) + aead, _ = chacha20poly1305.New(key[:]) + var ciphertext [MLKEMCiphertextSize]byte + _, err = aead.Open(ciphertext[:0], ZeroNonce[:], msg.MLKEM[:], hash[:]) + if err != nil { + return nil + } - var timestamp tai64n.Timestamp + mixHash(&hash, &hash, msg.MLKEM[:]) - handshake.mutex.RLock() + // post-quantum decapsulation + scheme := kyber1024.Scheme() + sk, err := scheme.UnmarshalBinaryPrivateKey(device.staticIdentity.mlkemPrivateKey[:]) + if err != nil { + return nil + } + + mlkemSecret, err := scheme.Decapsulate(sk, ciphertext[:]) + if err != nil { + return nil + } + // mix classic (ss) and post-quantum secret (mlkemSecret) into a single secret + var combinedSecret [blake2s.Size]byte + KDF2(&combinedSecret, nil, ss[:], mlkemSecret) + + // main chainKey is now updated with the combined secret + KDF2(&chainKey, &key, chainKey[:], combinedSecret[:]) + + // verify identity + var timestamp tai64n.Timestamp + handshake.mutex.RLock() if isZero(handshake.precomputedStaticStatic[:]) { handshake.mutex.RUnlock() return nil @@ -424,16 +452,17 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { chainKey[:], handshake.precomputedStaticStatic[:], ) + handshake.mutex.RUnlock() + aead, _ = chacha20poly1305.New(key[:]) _, err = aead.Open(timestamp[:0], ZeroNonce[:], msg.Timestamp[:], hash[:]) if err != nil { - handshake.mutex.RUnlock() return nil } mixHash(&hash, &hash, msg.Timestamp[:]) // protect against replay & flood - + handshake.mutex.RLock() replay := !timestamp.After(handshake.lastTimestamp) flood := time.Since(handshake.lastInitiationConsumption) <= HandshakeInitationRate handshake.mutex.RUnlock() @@ -447,9 +476,7 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { } // update handshake state - handshake.mutex.Lock() - handshake.hash = hash handshake.chainKey = chainKey handshake.remoteIndex = msg.Sender @@ -462,7 +489,6 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { handshake.lastInitiationConsumption = now } handshake.state = handshakeInitiationConsumed - handshake.mutex.Unlock() setZero(hash[:]) From 68be555e8635a94ee4f184d19e643160ef173ce1 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Tue, 23 Sep 2025 20:31:56 -0300 Subject: [PATCH 20/34] feat(device): Add function for quantum key generation Signed-off-by: Mateus Franco --- device/quantum-keys.go | 15 +++++++++++++++ 1 file changed, 15 insertions(+) create mode 100644 device/quantum-keys.go diff --git a/device/quantum-keys.go b/device/quantum-keys.go new file mode 100644 index 000000000..9c60e3eb0 --- /dev/null +++ b/device/quantum-keys.go @@ -0,0 +1,15 @@ +package device + +import ( + "github.com/cloudflare/circl/kem/kyber/kyber1024" +) + +func GenerateQuantumKeyPair() (pub []byte, priv []byte, err error) { + pk, sk, err := kyber1024.Scheme().GenerateKeyPair() + if err != nil { + return nil, nil, err + } + pub, _ = pk.MarshalBinary() + priv, _ = sk.MarshalBinary() + return pub, priv, nil +} From 87626c1f4acb7240498bd7a433a4feba00651052 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Tue, 23 Sep 2025 20:39:27 -0300 Subject: [PATCH 21/34] feat(device): Enhance UAPI handling for ML-KEM keys Signed-off-by: Mateus Franco --- device/uapi.go | 22 +++++++++++++++++++++- 1 file changed, 21 insertions(+), 1 deletion(-) diff --git a/device/uapi.go b/device/uapi.go index cc69488b4..d3ebb9b68 100644 --- a/device/uapi.go +++ b/device/uapi.go @@ -171,7 +171,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) { // Load/create the peer we are now configuring. err := device.handlePublicKeyLine(peer, value) if err != nil { - return err + return ipcErrorf(ipc.IpcErrorInvalid, "failed to load MLKEM public key: %w", err) } continue } @@ -240,6 +240,16 @@ func (device *Device) handleDeviceLine(key, value string) error { device.log.Verbosef("UAPI: Removing all peers") device.RemoveAllPeers() + case "mlkem_private_key": + var mlkemPrivateKey MLKEMPrivateKey + err := loadExactHex(mlkemPrivateKey[:], value) + if err != nil { + return err + } + device.staticIdentity.Lock() + device.staticIdentity.mlkemPrivateKey = mlkemPrivateKey + device.staticIdentity.Unlock() + default: return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key) } @@ -397,6 +407,16 @@ func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error return ipcErrorf(ipc.IpcErrorInvalid, "invalid protocol version: %v", value) } + case "mlkem_public_key": + device.log.Verbosef("%v - UAPI: Updating mlkem_public_key", peer.Peer) + peer.handshake.mutex.Lock() + err := loadExactHex(peer.handshake.remoteMLKEMStatic[:], value) + if err != nil { + peer.handshake.mutex.Unlock() + return err + } + peer.handshake.mutex.Unlock() + default: return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI peer key: %v", key) } From 0f31fc39943b21a048a4d01932d484442f0dbe88 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Tue, 23 Sep 2025 21:05:27 -0300 Subject: [PATCH 22/34] feat(device): Update KDF2 calls to include dummy parameter for combined secret derivation Signed-off-by: Mateus Franco --- device/noise-protocol.go | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 22c579184..ac68e34c7 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -326,7 +326,8 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e // mix classic (ss) and post-quantum secret (mlkemSecret) into a single secret var combinedSecret [blake2s.Size]byte - KDF2(&combinedSecret, nil, ss[:], mlkemSecret) + var dummy [blake2s.Size]byte + KDF2(&combinedSecret, &dummy, ss[:], mlkemSecret) // from this point on, use the combined secret to feed the Noise KDF KDF2( @@ -434,7 +435,8 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { // mix classic (ss) and post-quantum secret (mlkemSecret) into a single secret var combinedSecret [blake2s.Size]byte - KDF2(&combinedSecret, nil, ss[:], mlkemSecret) + var dummy [blake2s.Size]byte + KDF2(&combinedSecret, &dummy, ss[:], mlkemSecret) // main chainKey is now updated with the combined secret KDF2(&chainKey, &key, chainKey[:], combinedSecret[:]) From a05eabe97e829adc8b3e99a1e17d559c46b5ee84 Mon Sep 17 00:00:00 2001 From: Eruel6 Date: Mon, 29 Sep 2025 17:31:44 -0300 Subject: [PATCH 23/34] fix old tests and add new ones Signed-off-by: Mateus Franco --- device/device_test.go | 155 +++++++++++++++++++++++++++++++++++++++ device/noise-protocol.go | 49 ++++--------- 2 files changed, 171 insertions(+), 33 deletions(-) diff --git a/device/device_test.go b/device/device_test.go index 0091e2052..6512d15c6 100644 --- a/device/device_test.go +++ b/device/device_test.go @@ -20,6 +20,8 @@ import ( "testing" "time" + "github.com/cloudflare/circl/kem/kyber/kyber1024" + "golang.zx2c4.com/wireguard/conn" "golang.zx2c4.com/wireguard/conn/bindtest" "golang.zx2c4.com/wireguard/tun" @@ -189,6 +191,9 @@ func genTestPair(tb testing.TB, realSocket bool) (pair testPair) { // The device is ready. Close it when the test completes. tb.Cleanup(p.dev.Close) } + + installMLKEMKeys(tb, &pair) + return } @@ -473,4 +478,154 @@ func TestBatchSize(t *testing.T) { if want, got := 128, d.BatchSize(); got != want { t.Errorf("expected batch size %d, got %d", want, got) } +} + // instala ML-KEM nos dois lados do par de teste +func installMLKEMKeys(t testing.TB, pair *testPair) { + t.Helper() + scheme := kyber1024.Scheme() + pk0, sk0, err := scheme.GenerateKeyPair() + if err != nil { t.Fatal(err) } + pk0b, _ := pk0.MarshalBinary() + sk0b, _ := sk0.MarshalBinary() + pk1, sk1, err := scheme.GenerateKeyPair() + if err != nil { t.Fatal(err) } + pk1b, _ := pk1.MarshalBinary() + sk1b, _ := sk1.MarshalBinary() + + if err := pair[0].dev.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(sk0b))); err != nil { t.Fatal(err) } + if err := pair[1].dev.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(sk1b))); err != nil { t.Fatal(err) } + var pub0, pub1 NoisePublicKey + for k := range pair[0].dev.peers.keyMap { pub0 = k; break } + for k := range pair[1].dev.peers.keyMap { pub1 = k; break } + cfgPeer0 := uapiCfg( + "public_key", hex.EncodeToString(pub0[:]), + "mlkem_public_key", hex.EncodeToString(pk1b), + ) + cfgPeer1 := uapiCfg( + "public_key", hex.EncodeToString(pub1[:]), + "mlkem_public_key", hex.EncodeToString(pk0b), + ) + if err := pair[0].dev.IpcSet(cfgPeer0); err != nil { t.Fatal(err) } + if err := pair[1].dev.IpcSet(cfgPeer1); err != nil { t.Fatal(err) } +} + +// 1) Geração de chaves ML-KEM +func TestMLKEMKeyGeneration(t *testing.T) { + pub, priv, err := GenerateQuantumKeyPair() + if err != nil { t.Fatal(err) } + + scheme := kyber1024.Scheme() + if len(pub) != scheme.PublicKeySize() { + t.Fatalf("pub size mismatch: got %d, want %d", len(pub), scheme.PublicKeySize()) + } + if len(priv) != scheme.PrivateKeySize() { + t.Fatalf("priv size mismatch: got %d, want %d", len(priv), scheme.PrivateKeySize()) + } + + // Unmarshal deve funcionar + if _, err := scheme.UnmarshalBinaryPublicKey(pub); err != nil { t.Fatal(err) } + if _, err := scheme.UnmarshalBinaryPrivateKey(priv); err != nil { t.Fatal(err) } +} + +// 2) Encaps/Decaps produz o mesmo segredo +func TestMLKEMEncapDecap(t *testing.T) { + scheme := kyber1024.Scheme() + pk, sk, err := scheme.GenerateKeyPair() + if err != nil { t.Fatal(err) } + + ct, ssEnc, err := scheme.Encapsulate(pk) + if err != nil { t.Fatal(err) } + ssDec, err := scheme.Decapsulate(sk, ct) + if err != nil { t.Fatal(err) } + + if !bytes.Equal(ssEnc, ssDec) { + t.Fatal("ML-KEM shared secrets differ") + } +} + +// 3) Handshake completo com ML-KEM integrado (derivação de chaves de sessão) +func TestNoiseHandshakeWithMLKEM(t *testing.T) { + // cria dois devices com tun/bind reais o suficiente p/ handshake em memória + skA, _ := newPrivateKey() + skB, _ := newPrivateKey() + + tunA := tuntest.NewChannelTUN() + tunB := tuntest.NewChannelTUN() + + devA := NewDevice(tunA.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + devB := NewDevice(tunB.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + + defer devA.Close() + defer devB.Close() + + // configura Noise keys + if err := devA.SetPrivateKey(skA); err != nil { t.Fatal(err) } + if err := devB.SetPrivateKey(skB); err != nil { t.Fatal(err) } + + // cria peers (um apontando pro outro) + peerB, err := devA.NewPeer(skB.publicKey()) + if err != nil { t.Fatal(err) } + peerA, err := devB.NewPeer(skA.publicKey()) + if err != nil { t.Fatal(err) } + + // gera chaves ML-KEM e injeta + scheme := kyber1024.Scheme() + pkA, skAkem, _ := scheme.GenerateKeyPair() + pkB, skBkem, _ := scheme.GenerateKeyPair() + pkAb, _ := pkA.MarshalBinary() + pkBb, _ := pkB.MarshalBinary() + skAb, _ := skAkem.MarshalBinary() + skBb, _ := skBkem.MarshalBinary() + + // via UAPI para usar o mesmo caminho de produção + if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { t.Fatal(err) } + if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { t.Fatal(err) } + + // amarra a ML-KEM pub do remoto em cada peer + if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerB.handshake.remoteStatic[:]), + "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { t.Fatal(err) } + if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), + "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { t.Fatal(err) } + + // inicia peers (como no teste de Noise clássico) + peerA.Start() + peerB.Start() + + // 3.1 initiation + msg1, err := devA.CreateMessageInitiation(peerB) + if err != nil { t.Fatal(err) } + if p := devB.ConsumeMessageInitiation(msg1); p == nil { + t.Fatal("handshake failed at initiation (ML-KEM)") + } + + // 3.2 response + msg2, err := devB.CreateMessageResponse(peerA) + if err != nil { t.Fatal(err) } + if p := devA.ConsumeMessageResponse(msg2); p == nil { + t.Fatal("handshake failed at response (ML-KEM)") + } + + // 3.3 deriva chaves de sessão em ambos os lados + if err := peerA.BeginSymmetricSession(); err != nil { t.Fatal(err) } + if err := peerB.BeginSymmetricSession(); err != nil { t.Fatal(err) } + + // 3.4 valida criptografia / decriptação nas duas direções + keyA := peerA.keypairs.next.Load() + keyB := peerB.keypairs.current + + msg := []byte("pqc wireguard ok") + var nonce [12]byte + + // A -> B + out := keyA.send.Seal(nil, nonce[:], msg, nil) + plain, err := keyB.receive.Open(nil, nonce[:], out, nil) + if err != nil { t.Fatal(err) } + if !bytes.Equal(plain, msg) { t.Fatal("A->B decrypt mismatch") } + + // B -> A + out = keyB.send.Seal(nil, nonce[:], msg, nil) + plain, err = keyA.receive.Open(nil, nonce[:], out, nil) + if err != nil { t.Fatal(err) } + if !bytes.Equal(plain, msg) { t.Fatal("B->A decrypt mismatch") } + } diff --git a/device/noise-protocol.go b/device/noise-protocol.go index ac68e34c7..e02f60990 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -280,7 +280,6 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e handshake.mutex.Lock() defer handshake.mutex.Unlock() - // create ephemeral key var err error handshake.hash = InitialHash handshake.chainKey = InitialChainKey @@ -299,63 +298,47 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e handshake.mixKey(msg.Ephemeral[:]) handshake.mixHash(msg.Ephemeral[:]) - // post-quantum encapsulation + ss, err := handshake.localEphemeral.sharedSecret(handshake.remoteStatic) + if err != nil { + return nil, err + } + + var key [chacha20poly1305.KeySize]byte + var tempChainKey [blake2s.Size]byte + KDF2(&tempChainKey, &key, handshake.chainKey[:], ss[:]) + + aead, _ := chacha20poly1305.New(key[:]) + aead.Seal(msg.Static[:0], ZeroNonce[:], device.staticIdentity.publicKey[:], handshake.hash[:]) + handshake.mixHash(msg.Static[:]) + scheme := kyber1024.Scheme() pk, err := scheme.UnmarshalBinaryPublicKey(handshake.remoteMLKEMStatic[:]) if err != nil { return nil, err } - ciphertext, mlkemSecret, err := scheme.Encapsulate(pk) if err != nil { return nil, err } - // encrypy KEM ciphertext and mix into handshake hash - var key [chacha20poly1305.KeySize]byte KDF1(&key, handshake.chainKey[:], []byte("pqc-ciphertext-key")) - aead, _ := chacha20poly1305.New(key[:]) + aead, _ = chacha20poly1305.New(key[:]) aead.Seal(msg.MLKEM[:0], ZeroNonce[:], ciphertext, handshake.hash[:]) handshake.mixHash(msg.MLKEM[:]) - // encrypt static key - ss, err := handshake.localEphemeral.sharedSecret(handshake.remoteStatic) - if err != nil { - return nil, err - } - - // mix classic (ss) and post-quantum secret (mlkemSecret) into a single secret var combinedSecret [blake2s.Size]byte var dummy [blake2s.Size]byte KDF2(&combinedSecret, &dummy, ss[:], mlkemSecret) + KDF2(&handshake.chainKey, &key, handshake.chainKey[:], combinedSecret[:]) - // from this point on, use the combined secret to feed the Noise KDF - KDF2( - &handshake.chainKey, - &key, - handshake.chainKey[:], - combinedSecret[:], - ) - - aead, _ = chacha20poly1305.New(key[:]) - aead.Seal(msg.Static[:0], ZeroNonce[:], device.staticIdentity.publicKey[:], handshake.hash[:]) - handshake.mixHash(msg.Static[:]) - - // encrypt timestamp if isZero(handshake.precomputedStaticStatic[:]) { return nil, errInvalidPublicKey } - KDF2( - &handshake.chainKey, - &key, - handshake.chainKey[:], - handshake.precomputedStaticStatic[:], - ) + KDF2(&handshake.chainKey, &key, handshake.chainKey[:], handshake.precomputedStaticStatic[:]) timestamp := tai64n.Now() aead, _ = chacha20poly1305.New(key[:]) aead.Seal(msg.Timestamp[:0], ZeroNonce[:], timestamp[:], handshake.hash[:]) - // assign index device.indexTable.Delete(handshake.localIndex) msg.Sender, err = device.indexTable.NewIndexForHandshake(peer, handshake) if err != nil { From 2f5176259ebceec70dd17fdbec57d0d78eaaea29 Mon Sep 17 00:00:00 2001 From: Eruel6 Date: Tue, 30 Sep 2025 11:28:06 -0300 Subject: [PATCH 24/34] feat: add benchmark test file Signed-off-by: Mateus Franco --- device/mlkem_bench_test.go | 159 +++++++++++++++++++++++++++++++++++++ 1 file changed, 159 insertions(+) create mode 100644 device/mlkem_bench_test.go diff --git a/device/mlkem_bench_test.go b/device/mlkem_bench_test.go new file mode 100644 index 000000000..405fa0e3e --- /dev/null +++ b/device/mlkem_bench_test.go @@ -0,0 +1,159 @@ +package device + +import ( + "bytes" + "encoding/hex" + "testing" + "time" + "golang.zx2c4.com/wireguard/tai64n" + + "github.com/cloudflare/circl/kem/kyber/kyber1024" + "golang.zx2c4.com/wireguard/conn" + "golang.zx2c4.com/wireguard/tun/tuntest" +) + +func BenchmarkKyberEncapsulate(b *testing.B) { + scheme := kyber1024.Scheme() + pk, _, err := scheme.GenerateKeyPair() + if err != nil { b.Fatal(err) } + b.ReportAllocs() + for i := 0; i < b.N; i++ { + _, _, err := scheme.Encapsulate(pk) + if err != nil { b.Fatal(err) } + } +} + +func BenchmarkKyberDecapsulate(b *testing.B) { + scheme := kyber1024.Scheme() + pk, sk, err := scheme.GenerateKeyPair() + if err != nil { b.Fatal(err) } + ct, _, err := scheme.Encapsulate(pk) // ct fixo é OK para decap + if err != nil { b.Fatal(err) } + b.ReportAllocs() + for i := 0; i < b.N; i++ { + _, err := scheme.Decapsulate(sk, ct) + if err != nil { b.Fatal(err) } + } +} + +// Handshake fim-a-fim (sem rede), com ML-KEM integrado +func BenchmarkHandshakeWithMLKEM(b *testing.B) { + // devs e peers em memória + skA, _ := newPrivateKey() + skB, _ := newPrivateKey() + tunA := tuntest.NewChannelTUN() + tunB := tuntest.NewChannelTUN() + devA := NewDevice(tunA.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + devB := NewDevice(tunB.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + defer devA.Close() + defer devB.Close() + + if err := devA.SetPrivateKey(skA); err != nil { b.Fatal(err) } + if err := devB.SetPrivateKey(skB); err != nil { b.Fatal(err) } + + peerB, err := devA.NewPeer(skB.publicKey()) + if err != nil { b.Fatal(err) } + peerA, err := devB.NewPeer(skA.publicKey()) + if err != nil { b.Fatal(err) } + + // injeta ML-KEM via UAPI (como em produção) + scheme := kyber1024.Scheme() + pkA, skAkem, _ := scheme.GenerateKeyPair() + pkB, skBkem, _ := scheme.GenerateKeyPair() + pkAb, _ := pkA.MarshalBinary() + pkBb, _ := pkB.MarshalBinary() + skAb, _ := skAkem.MarshalBinary() + skBb, _ := skBkem.MarshalBinary() + + if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { b.Fatal(err) } + if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { b.Fatal(err) } + if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerB.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { b.Fatal(err) } + if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { b.Fatal(err) } + + peerA.Start() + peerB.Start() + + // >>> helper local: relaxa o rate-limit do receptor (devB/peerA) + relaxFlood := func() { + peerA.handshake.mutex.Lock() + // evita "flood" + peerA.handshake.lastInitiationConsumption = time.Now().Add(-10 * time.Second) + // evita "replay" + peerA.handshake.lastTimestamp = tai64n.Timestamp{} // zero + peerA.handshake.mutex.Unlock() +} + + // warmup + relaxFlood() + msg1, err := devA.CreateMessageInitiation(peerB); if err != nil { b.Fatal(err) } + if p := devB.ConsumeMessageInitiation(msg1); p == nil { b.Fatal("initiation fail (warmup)") } + msg2, err := devB.CreateMessageResponse(peerA); if err != nil { b.Fatal(err) } + if p := devA.ConsumeMessageResponse(msg2); p == nil { b.Fatal("response fail (warmup)") } + if err := peerA.BeginSymmetricSession(); err != nil { b.Fatal(err) } + if err := peerB.BeginSymmetricSession(); err != nil { b.Fatal(err) } + + b.ReportAllocs() + b.ResetTimer() + + for i := 0; i < b.N; i++ { + relaxFlood() // <<< chama antes de cada initiation + msg1, err := devA.CreateMessageInitiation(peerB); if err != nil { b.Fatal(err) } + if p := devB.ConsumeMessageInitiation(msg1); p == nil { b.Fatal("initiation fail") } + msg2, err := devB.CreateMessageResponse(peerA); if err != nil { b.Fatal(err) } + if p := devA.ConsumeMessageResponse(msg2); p == nil { b.Fatal("response fail") } + if err := peerA.BeginSymmetricSession(); err != nil { b.Fatal(err) } + if err := peerB.BeginSymmetricSession(); err != nil { b.Fatal(err) } + } +} + +// (opcional) mede a cifra/decifra de um payload curto com as chaves da sessão +func BenchmarkDataPlaneAEAD(b *testing.B) { + // reaproveita o setup do handshake acima + skA, _ := newPrivateKey() + skB, _ := newPrivateKey() + tunA := tuntest.NewChannelTUN() + tunB := tuntest.NewChannelTUN() + devA := NewDevice(tunA.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + devB := NewDevice(tunB.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + defer devA.Close() + defer devB.Close() + devA.SetPrivateKey(skA); devB.SetPrivateKey(skB) + peerB, _ := devA.NewPeer(skB.publicKey()) + peerA, _ := devB.NewPeer(skA.publicKey()) + + // ML-KEM + scheme := kyber1024.Scheme() + pkA, skAkem, _ := scheme.GenerateKeyPair() + pkB, skBkem, _ := scheme.GenerateKeyPair() + pkAb, _ := pkA.MarshalBinary() + pkBb, _ := pkB.MarshalBinary() + skAb, _ := skAkem.MarshalBinary() + skBb, _ := skBkem.MarshalBinary() + devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))) + devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))) + devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerB.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkBb))) + devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkAb))) + peerA.Start(); peerB.Start() + + // establish one session + msg1, _ := devA.CreateMessageInitiation(peerB) + devB.ConsumeMessageInitiation(msg1) + msg2, _ := devB.CreateMessageResponse(peerA) + devA.ConsumeMessageResponse(msg2) + peerA.BeginSymmetricSession() + peerB.BeginSymmetricSession() + + keyA := peerA.keypairs.next.Load() + keyB := peerB.keypairs.current + msg := bytes.Repeat([]byte{0x42}, 128) // 128B + var nonce [12]byte + + b.ReportAllocs() + b.SetBytes(int64(len(msg))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + out := keyA.send.Seal(nil, nonce[:], msg, nil) + _, err := keyB.receive.Open(nil, nonce[:], out, nil) + if err != nil { b.Fatal(err) } + } +} From 92c9c80e4ed6b5c7a21c5c00b40228cc263610e9 Mon Sep 17 00:00:00 2001 From: Eruel6 Date: Thu, 16 Oct 2025 10:57:40 -0300 Subject: [PATCH 25/34] fix: add new tests to utlize hybrid handshake Signed-off-by: Mateus Franco --- device/device_test.go | 21 +-- device/mlkem_bench_test.go | 292 ++++++++++++++++++++++++++++++------- 2 files changed, 242 insertions(+), 71 deletions(-) diff --git a/device/device_test.go b/device/device_test.go index 6512d15c6..42da6e1b0 100644 --- a/device/device_test.go +++ b/device/device_test.go @@ -410,7 +410,6 @@ func goroutineLeakCheck(t *testing.T) { if t.Failed() { return } - // Give goroutines time to exit, if they need it. for i := 0; i < 10000; i++ { if runtime.NumGoroutine() <= startGoroutines { return @@ -479,7 +478,7 @@ func TestBatchSize(t *testing.T) { t.Errorf("expected batch size %d, got %d", want, got) } } - // instala ML-KEM nos dois lados do par de teste + func installMLKEMKeys(t testing.TB, pair *testPair) { t.Helper() scheme := kyber1024.Scheme() @@ -509,7 +508,6 @@ func installMLKEMKeys(t testing.TB, pair *testPair) { if err := pair[1].dev.IpcSet(cfgPeer1); err != nil { t.Fatal(err) } } -// 1) Geração de chaves ML-KEM func TestMLKEMKeyGeneration(t *testing.T) { pub, priv, err := GenerateQuantumKeyPair() if err != nil { t.Fatal(err) } @@ -522,12 +520,11 @@ func TestMLKEMKeyGeneration(t *testing.T) { t.Fatalf("priv size mismatch: got %d, want %d", len(priv), scheme.PrivateKeySize()) } - // Unmarshal deve funcionar if _, err := scheme.UnmarshalBinaryPublicKey(pub); err != nil { t.Fatal(err) } if _, err := scheme.UnmarshalBinaryPrivateKey(priv); err != nil { t.Fatal(err) } } -// 2) Encaps/Decaps produz o mesmo segredo + func TestMLKEMEncapDecap(t *testing.T) { scheme := kyber1024.Scheme() pk, sk, err := scheme.GenerateKeyPair() @@ -543,9 +540,7 @@ func TestMLKEMEncapDecap(t *testing.T) { } } -// 3) Handshake completo com ML-KEM integrado (derivação de chaves de sessão) func TestNoiseHandshakeWithMLKEM(t *testing.T) { - // cria dois devices com tun/bind reais o suficiente p/ handshake em memória skA, _ := newPrivateKey() skB, _ := newPrivateKey() @@ -558,17 +553,14 @@ func TestNoiseHandshakeWithMLKEM(t *testing.T) { defer devA.Close() defer devB.Close() - // configura Noise keys if err := devA.SetPrivateKey(skA); err != nil { t.Fatal(err) } if err := devB.SetPrivateKey(skB); err != nil { t.Fatal(err) } - // cria peers (um apontando pro outro) peerB, err := devA.NewPeer(skB.publicKey()) if err != nil { t.Fatal(err) } peerA, err := devB.NewPeer(skA.publicKey()) if err != nil { t.Fatal(err) } - // gera chaves ML-KEM e injeta scheme := kyber1024.Scheme() pkA, skAkem, _ := scheme.GenerateKeyPair() pkB, skBkem, _ := scheme.GenerateKeyPair() @@ -577,52 +569,43 @@ func TestNoiseHandshakeWithMLKEM(t *testing.T) { skAb, _ := skAkem.MarshalBinary() skBb, _ := skBkem.MarshalBinary() - // via UAPI para usar o mesmo caminho de produção if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { t.Fatal(err) } if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { t.Fatal(err) } - // amarra a ML-KEM pub do remoto em cada peer if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerB.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { t.Fatal(err) } if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { t.Fatal(err) } - // inicia peers (como no teste de Noise clássico) peerA.Start() peerB.Start() - // 3.1 initiation msg1, err := devA.CreateMessageInitiation(peerB) if err != nil { t.Fatal(err) } if p := devB.ConsumeMessageInitiation(msg1); p == nil { t.Fatal("handshake failed at initiation (ML-KEM)") } - // 3.2 response msg2, err := devB.CreateMessageResponse(peerA) if err != nil { t.Fatal(err) } if p := devA.ConsumeMessageResponse(msg2); p == nil { t.Fatal("handshake failed at response (ML-KEM)") } - // 3.3 deriva chaves de sessão em ambos os lados if err := peerA.BeginSymmetricSession(); err != nil { t.Fatal(err) } if err := peerB.BeginSymmetricSession(); err != nil { t.Fatal(err) } - // 3.4 valida criptografia / decriptação nas duas direções keyA := peerA.keypairs.next.Load() keyB := peerB.keypairs.current msg := []byte("pqc wireguard ok") var nonce [12]byte - // A -> B out := keyA.send.Seal(nil, nonce[:], msg, nil) plain, err := keyB.receive.Open(nil, nonce[:], out, nil) if err != nil { t.Fatal(err) } if !bytes.Equal(plain, msg) { t.Fatal("A->B decrypt mismatch") } - // B -> A out = keyB.send.Seal(nil, nonce[:], msg, nil) plain, err = keyA.receive.Open(nil, nonce[:], out, nil) if err != nil { t.Fatal(err) } diff --git a/device/mlkem_bench_test.go b/device/mlkem_bench_test.go index 405fa0e3e..349f205e9 100644 --- a/device/mlkem_bench_test.go +++ b/device/mlkem_bench_test.go @@ -4,8 +4,9 @@ import ( "bytes" "encoding/hex" "testing" - "time" - "golang.zx2c4.com/wireguard/tai64n" + "time" + + "golang.zx2c4.com/wireguard/tai64n" "github.com/cloudflare/circl/kem/kyber/kyber1024" "golang.zx2c4.com/wireguard/conn" @@ -15,30 +16,38 @@ import ( func BenchmarkKyberEncapsulate(b *testing.B) { scheme := kyber1024.Scheme() pk, _, err := scheme.GenerateKeyPair() - if err != nil { b.Fatal(err) } + if err != nil { + b.Fatal(err) + } b.ReportAllocs() for i := 0; i < b.N; i++ { _, _, err := scheme.Encapsulate(pk) - if err != nil { b.Fatal(err) } + if err != nil { + b.Fatal(err) + } } } func BenchmarkKyberDecapsulate(b *testing.B) { scheme := kyber1024.Scheme() pk, sk, err := scheme.GenerateKeyPair() - if err != nil { b.Fatal(err) } - ct, _, err := scheme.Encapsulate(pk) // ct fixo é OK para decap - if err != nil { b.Fatal(err) } + if err != nil { + b.Fatal(err) + } + ct, _, err := scheme.Encapsulate(pk) + if err != nil { + b.Fatal(err) + } b.ReportAllocs() for i := 0; i < b.N; i++ { _, err := scheme.Decapsulate(sk, ct) - if err != nil { b.Fatal(err) } + if err != nil { + b.Fatal(err) + } } } -// Handshake fim-a-fim (sem rede), com ML-KEM integrado func BenchmarkHandshakeWithMLKEM(b *testing.B) { - // devs e peers em memória skA, _ := newPrivateKey() skB, _ := newPrivateKey() tunA := tuntest.NewChannelTUN() @@ -48,15 +57,22 @@ func BenchmarkHandshakeWithMLKEM(b *testing.B) { defer devA.Close() defer devB.Close() - if err := devA.SetPrivateKey(skA); err != nil { b.Fatal(err) } - if err := devB.SetPrivateKey(skB); err != nil { b.Fatal(err) } + if err := devA.SetPrivateKey(skA); err != nil { + b.Fatal(err) + } + if err := devB.SetPrivateKey(skB); err != nil { + b.Fatal(err) + } peerB, err := devA.NewPeer(skB.publicKey()) - if err != nil { b.Fatal(err) } + if err != nil { + b.Fatal(err) + } peerA, err := devB.NewPeer(skA.publicKey()) - if err != nil { b.Fatal(err) } + if err != nil { + b.Fatal(err) + } - // injeta ML-KEM via UAPI (como em produção) scheme := kyber1024.Scheme() pkA, skAkem, _ := scheme.GenerateKeyPair() pkB, skBkem, _ := scheme.GenerateKeyPair() @@ -65,50 +81,151 @@ func BenchmarkHandshakeWithMLKEM(b *testing.B) { skAb, _ := skAkem.MarshalBinary() skBb, _ := skBkem.MarshalBinary() - if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { b.Fatal(err) } - if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { b.Fatal(err) } - if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerB.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { b.Fatal(err) } - if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { b.Fatal(err) } - - peerA.Start() - peerB.Start() + if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { + b.Fatal(err) + } + if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerB.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { + b.Fatal(err) + } - // >>> helper local: relaxa o rate-limit do receptor (devB/peerA) relaxFlood := func() { - peerA.handshake.mutex.Lock() - // evita "flood" - peerA.handshake.lastInitiationConsumption = time.Now().Add(-10 * time.Second) - // evita "replay" - peerA.handshake.lastTimestamp = tai64n.Timestamp{} // zero - peerA.handshake.mutex.Unlock() + peerA.handshake.mutex.Lock() + peerA.handshake.lastInitiationConsumption = time.Now().Add(-10 * time.Second) + peerA.handshake.lastTimestamp = tai64n.Timestamp{} + peerA.handshake.mutex.Unlock() + } + relaxFlood() + msg1, err := devA.CreateMessageInitiation(peerB) + if err != nil { + b.Fatal(err) + } + if p := devB.ConsumeMessageInitiation(msg1); p == nil { + b.Fatal("initiation fail (warmup)") + } + msg2, err := devB.CreateMessageResponse(peerA) + if err != nil { + b.Fatal(err) + } + if p := devA.ConsumeMessageResponse(msg2); p == nil { + b.Fatal("response fail (warmup)") + } + if err := peerA.BeginSymmetricSession(); err != nil { + b.Fatal(err) + } + if err := peerB.BeginSymmetricSession(); err != nil { + b.Fatal(err) + } + + b.ReportAllocs() + b.ResetTimer() + + for i := 0; i < b.N; i++ { + relaxFlood() + msg1, err := devA.CreateMessageInitiation(peerB) + if err != nil { + b.Fatal(err) + } + if p := devB.ConsumeMessageInitiation(msg1); p == nil { + b.Fatal("initiation fail") + } + msg2, err := devB.CreateMessageResponse(peerA) + if err != nil { + b.Fatal(err) + } + if p := devA.ConsumeMessageResponse(msg2); p == nil { + b.Fatal("response fail") + } + if err := peerA.BeginSymmetricSession(); err != nil { + b.Fatal(err) + } + if err := peerB.BeginSymmetricSession(); err != nil { + b.Fatal(err) + } + } } - // warmup - relaxFlood() - msg1, err := devA.CreateMessageInitiation(peerB); if err != nil { b.Fatal(err) } - if p := devB.ConsumeMessageInitiation(msg1); p == nil { b.Fatal("initiation fail (warmup)") } - msg2, err := devB.CreateMessageResponse(peerA); if err != nil { b.Fatal(err) } - if p := devA.ConsumeMessageResponse(msg2); p == nil { b.Fatal("response fail (warmup)") } - if err := peerA.BeginSymmetricSession(); err != nil { b.Fatal(err) } - if err := peerB.BeginSymmetricSession(); err != nil { b.Fatal(err) } +func BenchmarkHandshakeHybrid(b *testing.B) { + skA, _ := newPrivateKey() + skB, _ := newPrivateKey() + tunA := tuntest.NewChannelTUN() + tunB := tuntest.NewChannelTUN() + devA := NewDevice(tunA.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + devB := NewDevice(tunB.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + defer devA.Close() + defer devB.Close() + + if err := devA.IpcSet(uapiCfg("private_key", hex.EncodeToString(skA[:]))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("private_key", hex.EncodeToString(skB[:]))); err != nil { + b.Fatal(err) + } + + scheme := kyber1024.Scheme() + pkA, skAkem, _ := scheme.GenerateKeyPair() + pkB, skBkem, _ := scheme.GenerateKeyPair() + pkAb, _ := pkA.MarshalBinary() + pkBb, _ := pkB.MarshalBinary() + skAb, _ := skAkem.MarshalBinary() + skBb, _ := skBkem.MarshalBinary() + + if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { + b.Fatal(err) + } + + pkBNoise := skB.publicKey() + pkANoise := skA.publicKey() + if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(pkBNoise[:]), "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(pkANoise[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { + b.Fatal(err) + } + + peerB := devA.LookupPeer(pkBNoise) + peerA := devB.LookupPeer(pkANoise) + if peerA == nil || peerB == nil { + b.Fatal("peer lookup failed (check IpcSet order)") + } + + relax := func() { + peerA.handshake.mutex.Lock() + peerA.handshake.lastInitiationConsumption = time.Now().Add(-10 * time.Second) + peerA.handshake.lastTimestamp = tai64n.Timestamp{} + peerA.handshake.mutex.Unlock() + } + relax() + msg1, _ := devA.CreateMessageInitiation(peerB) + devB.ConsumeMessageInitiation(msg1) + msg2, _ := devB.CreateMessageResponse(peerA) + devA.ConsumeMessageResponse(msg2) + peerA.BeginSymmetricSession() + peerB.BeginSymmetricSession() b.ReportAllocs() b.ResetTimer() for i := 0; i < b.N; i++ { - relaxFlood() // <<< chama antes de cada initiation - msg1, err := devA.CreateMessageInitiation(peerB); if err != nil { b.Fatal(err) } - if p := devB.ConsumeMessageInitiation(msg1); p == nil { b.Fatal("initiation fail") } - msg2, err := devB.CreateMessageResponse(peerA); if err != nil { b.Fatal(err) } - if p := devA.ConsumeMessageResponse(msg2); p == nil { b.Fatal("response fail") } - if err := peerA.BeginSymmetricSession(); err != nil { b.Fatal(err) } - if err := peerB.BeginSymmetricSession(); err != nil { b.Fatal(err) } + relax() + msg1, _ := devA.CreateMessageInitiation(peerB) + devB.ConsumeMessageInitiation(msg1) + msg2, _ := devB.CreateMessageResponse(peerA) + devA.ConsumeMessageResponse(msg2) + peerA.BeginSymmetricSession() + peerB.BeginSymmetricSession() } } -// (opcional) mede a cifra/decifra de um payload curto com as chaves da sessão func BenchmarkDataPlaneAEAD(b *testing.B) { - // reaproveita o setup do handshake acima skA, _ := newPrivateKey() skB, _ := newPrivateKey() tunA := tuntest.NewChannelTUN() @@ -117,11 +234,11 @@ func BenchmarkDataPlaneAEAD(b *testing.B) { devB := NewDevice(tunB.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) defer devA.Close() defer devB.Close() - devA.SetPrivateKey(skA); devB.SetPrivateKey(skB) + devA.SetPrivateKey(skA) + devB.SetPrivateKey(skB) peerB, _ := devA.NewPeer(skB.publicKey()) peerA, _ := devB.NewPeer(skA.publicKey()) - // ML-KEM scheme := kyber1024.Scheme() pkA, skAkem, _ := scheme.GenerateKeyPair() pkB, skBkem, _ := scheme.GenerateKeyPair() @@ -133,9 +250,7 @@ func BenchmarkDataPlaneAEAD(b *testing.B) { devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))) devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerB.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkBb))) devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkAb))) - peerA.Start(); peerB.Start() - // establish one session msg1, _ := devA.CreateMessageInitiation(peerB) devB.ConsumeMessageInitiation(msg1) msg2, _ := devB.CreateMessageResponse(peerA) @@ -145,7 +260,78 @@ func BenchmarkDataPlaneAEAD(b *testing.B) { keyA := peerA.keypairs.next.Load() keyB := peerB.keypairs.current - msg := bytes.Repeat([]byte{0x42}, 128) // 128B + msg := bytes.Repeat([]byte{0x42}, 128) + var nonce [12]byte + + b.ReportAllocs() + b.SetBytes(int64(len(msg))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + out := keyA.send.Seal(nil, nonce[:], msg, nil) + _, err := keyB.receive.Open(nil, nonce[:], out, nil) + if err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkDataPlaneAEADHybrid(b *testing.B) { + skA, _ := newPrivateKey() + skB, _ := newPrivateKey() + tunA := tuntest.NewChannelTUN() + tunB := tuntest.NewChannelTUN() + devA := NewDevice(tunA.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + devB := NewDevice(tunB.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + defer devA.Close() + defer devB.Close() + + if err := devA.IpcSet(uapiCfg("private_key", hex.EncodeToString(skA[:]))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("private_key", hex.EncodeToString(skB[:]))); err != nil { + b.Fatal(err) + } + + scheme := kyber1024.Scheme() + pkA, skAkem, _ := scheme.GenerateKeyPair() + pkB, skBkem, _ := scheme.GenerateKeyPair() + pkAb, _ := pkA.MarshalBinary() + pkBb, _ := pkB.MarshalBinary() + skAb, _ := skAkem.MarshalBinary() + skBb, _ := skBkem.MarshalBinary() + + if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { + b.Fatal(err) + } + + pkBNoise := skB.publicKey() + pkANoise := skA.publicKey() + if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(pkBNoise[:]), "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(pkANoise[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { + b.Fatal(err) + } + + peerB := devA.LookupPeer(pkBNoise) + peerA := devB.LookupPeer(pkANoise) + if peerA == nil || peerB == nil { + b.Fatal("peer lookup failed (check IpcSet order)") + } + + msg1, _ := devA.CreateMessageInitiation(peerB) + devB.ConsumeMessageInitiation(msg1) + msg2, _ := devB.CreateMessageResponse(peerA) + devA.ConsumeMessageResponse(msg2) + peerA.BeginSymmetricSession() + peerB.BeginSymmetricSession() + + keyA := peerA.keypairs.next.Load() + keyB := peerB.keypairs.current + msg := bytes.Repeat([]byte{0x42}, 128) var nonce [12]byte b.ReportAllocs() @@ -154,6 +340,8 @@ func BenchmarkDataPlaneAEAD(b *testing.B) { for i := 0; i < b.N; i++ { out := keyA.send.Seal(nil, nonce[:], msg, nil) _, err := keyB.receive.Open(nil, nonce[:], out, nil) - if err != nil { b.Fatal(err) } + if err != nil { + b.Fatal(err) + } } } From d80d6d9115b10888558e84fd70917dae63ee218f Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Mon, 13 Oct 2025 19:59:30 -0300 Subject: [PATCH 26/34] feat(device): Add MLDSA key and signature types with appropriate sizes Signed-off-by: Mateus Franco --- device/noise-types.go | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/device/noise-types.go b/device/noise-types.go index 523ffa12c..afa53e784 100644 --- a/device/noise-types.go +++ b/device/noise-types.go @@ -87,3 +87,15 @@ type ( MLKEMPublicKey [MLKEMPublicKeySize]byte MLKEMPrivateKey [MLKEMPrivateKeySize]byte ) + +const ( + MLDSAPublicKeySize = 2592 + MLDSAPrivateKeySize = 4864 + MLDSASignatureSize = 4595 +) + +type ( + MLDSAPublicKey [MLDSAPublicKeySize]byte + MLDSAPrivateKey [MLDSAPrivateKeySize]byte + MLDSASignature [MLDSASignatureSize]byte +) From 26943f47dc92e91631c07df0aee437bd94dd21e1 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Mon, 13 Oct 2025 20:13:06 -0300 Subject: [PATCH 27/34] feat(device): Add MLDSA public key to Handshake struct and include private keys in Device struct Signed-off-by: Mateus Franco --- device/device.go | 2 ++ device/noise-protocol.go | 1 + 2 files changed, 3 insertions(+) diff --git a/device/device.go b/device/device.go index d889a9b81..6dbca4a67 100644 --- a/device/device.go +++ b/device/device.go @@ -53,6 +53,8 @@ type Device struct { publicKey NoisePublicKey mlkemPrivateKey MLKEMPrivateKey mlkemPublicKey MLKEMPublicKey + mldsaPrivateKey MLDSAPrivateKey + mldsaPublicKey MLDSAPublicKey } peers struct { diff --git a/device/noise-protocol.go b/device/noise-protocol.go index e02f60990..6aba6ed6f 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -223,6 +223,7 @@ type Handshake struct { remoteIndex uint32 // index for sending remoteStatic NoisePublicKey // long term key remoteMLKEMStatic MLKEMPublicKey // long term remote ML-KEM static public key + remoteMLDSAStatic MLDSAPublicKey // long term remote ML-DSA static public key remoteEphemeral NoisePublicKey // ephemeral public key precomputedStaticStatic [NoisePublicKeySize]byte // precomputed shared secret lastTimestamp tai64n.Timestamp From cada0b9197091c2e97139effc657e6c59cd7a6cf Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Mon, 13 Oct 2025 20:18:51 -0300 Subject: [PATCH 28/34] feat(device): Update handshake message sizes to include MLDSA signature Signed-off-by: Mateus Franco --- device/noise-protocol.go | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 6aba6ed6f..610fd85df 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -62,13 +62,13 @@ const ( ) const ( - MessageInitiationSize = 148 + (MLKEMCiphertextSize + poly1305.TagSize) // size of handshake initiation message - MessageResponseSize = 92 // size of response message - MessageCookieReplySize = 64 // size of cookie reply message - MessageTransportHeaderSize = 16 // size of data preceding content in transport message - MessageTransportSize = MessageTransportHeaderSize + poly1305.TagSize // size of empty transport - MessageKeepaliveSize = MessageTransportSize // size of keepalive - MessageHandshakeSize = MessageInitiationSize // size of largest handshake related message + MessageInitiationSize = 148 + (MLKEMCiphertextSize + poly1305.TagSize) + MLDSASignatureSize // size of handshake initiation message + MessageResponseSize = 92 // size of response message + MessageCookieReplySize = 64 // size of cookie reply message + MessageTransportHeaderSize = 16 // size of data preceding content in transport message + MessageTransportSize = MessageTransportHeaderSize + poly1305.TagSize // size of empty transport + MessageKeepaliveSize = MessageTransportSize // size of keepalive + MessageHandshakeSize = MessageInitiationSize // size of largest handshake related message ) const ( @@ -90,6 +90,7 @@ type MessageInitiation struct { Static [NoisePublicKeySize + poly1305.TagSize]byte MLKEM [MLKEMCiphertextSize + poly1305.TagSize]byte Timestamp [tai64n.TimestampSize + poly1305.TagSize]byte + Signature MLDSASignature MAC1 [blake2s.Size128]byte MAC2 [blake2s.Size128]byte } From 10896e6020f58427fa0c470ee69c02d90ccaceb9 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Mon, 13 Oct 2025 20:22:35 -0300 Subject: [PATCH 29/34] feat(device): Update MessageInitiation struct to include Signature in marshaling and unmarshaling Signed-off-by: Mateus Franco --- device/noise-protocol.go | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 610fd85df..adffc4560 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -132,8 +132,9 @@ func (msg *MessageInitiation) unmarshal(b []byte) error { copy(msg.Static[:], b[8+len(msg.Ephemeral):]) copy(msg.MLKEM[:], b[8+len(msg.Ephemeral)+len(msg.Static):]) copy(msg.Timestamp[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM):]) - copy(msg.MAC1[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp):]) - copy(msg.MAC2[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp)+len(msg.MAC1):]) + copy(msg.Signature[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp):]) + copy(msg.MAC1[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp)+len(msg.Signature):]) + copy(msg.MAC2[:], b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp)+len(msg.Signature)+len(msg.MAC1):]) return nil } @@ -149,8 +150,9 @@ func (msg *MessageInitiation) marshal(b []byte) error { copy(b[8+len(msg.Ephemeral):], msg.Static[:]) copy(b[8+len(msg.Ephemeral)+len(msg.Static):], msg.MLKEM[:]) copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM):], msg.Timestamp[:]) - copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp):], msg.MAC1[:]) - copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp)+len(msg.MAC1):], msg.MAC2[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp):], msg.Signature[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp)+len(msg.Signature):], msg.MAC1[:]) + copy(b[8+len(msg.Ephemeral)+len(msg.Static)+len(msg.MLKEM)+len(msg.Timestamp)+len(msg.Signature)+len(msg.MAC1):], msg.MAC2[:]) return nil } From fabb5d76b3e7635d9f5f3222cba7eefeb82978a3 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sun, 26 Oct 2025 11:59:55 -0300 Subject: [PATCH 30/34] feat(device): Implement MLDSA signature generation in CreateMessageInitiation Signed-off-by: Mateus Franco --- device/noise-protocol.go | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index adffc4560..8b8831d8f 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -13,6 +13,7 @@ import ( "time" "github.com/cloudflare/circl/kem/kyber/kyber1024" + "github.com/cloudflare/circl/sign/dilithium/mode5" "golang.org/x/crypto/blake2s" "golang.org/x/crypto/chacha20poly1305" "golang.org/x/crypto/poly1305" @@ -352,6 +353,21 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e handshake.mixHash(msg.Timestamp[:]) handshake.state = handshakeInitiationCreated + + signScheme := mode5.Scheme() + skSign, err := signScheme.UnmarshalBinaryPrivateKey(device.staticIdentity.mldsaPrivateKey[:]) + if err != nil { + return nil, err + } + + messageToSign := make([]byte, MessageInitiationSize) + if err := msg.marshal(messageToSign); err != nil { + return nil, err + } + + signature := signScheme.Sign(skSign, messageToSign[:MessageInitiationSize-blake2s.Size128*2-MLDSASignatureSize], nil) + copy(msg.Signature[:], signature) + return &msg, nil } From 4193bcc9910f9f50836b7a2e2608ee9251215e0c Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sun, 26 Oct 2025 12:03:55 -0300 Subject: [PATCH 31/34] feat(device): Add MLDSA signature verification in ConsumeMessageInitiation Signed-off-by: Mateus Franco --- device/noise-protocol.go | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 8b8831d8f..e0f7006f2 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -411,6 +411,21 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { return nil } + signScheme := mode5.Scheme() + pkSign, err := signScheme.UnmarshalBinaryPublicKey(peer.handshake.remoteMLDSAStatic[:]) + if err != nil { + return nil + } + + messageToCheck := make([]byte, MessageInitiationSize) + if err := msg.marshal(messageToCheck); err != nil { + return nil + } + + if !signScheme.Verify(pkSign, messageToCheck[:MessageInitiationSize-blake2s.Size128*2-MLDSASignatureSize], msg.Signature[:], nil) { + return nil + } + handshake := &peer.handshake // decrypt KEM ciphertext From 5a1d36e09241aade6cebfec1983a8998ef0a9206 Mon Sep 17 00:00:00 2001 From: Mateus Franco Date: Sun, 26 Oct 2025 12:18:00 -0300 Subject: [PATCH 32/34] feat(device): Add MLDSA key pair generation function Signed-off-by: Mateus Franco --- device/quantum-keys.go | 12 ++++++++++++ go.mod | 2 +- 2 files changed, 13 insertions(+), 1 deletion(-) diff --git a/device/quantum-keys.go b/device/quantum-keys.go index 9c60e3eb0..36575bada 100644 --- a/device/quantum-keys.go +++ b/device/quantum-keys.go @@ -2,6 +2,7 @@ package device import ( "github.com/cloudflare/circl/kem/kyber/kyber1024" + "github.com/cloudflare/circl/sign/dilithium/mode5" ) func GenerateQuantumKeyPair() (pub []byte, priv []byte, err error) { @@ -13,3 +14,14 @@ func GenerateQuantumKeyPair() (pub []byte, priv []byte, err error) { priv, _ = sk.MarshalBinary() return pub, priv, nil } + +func GenerateMLDSAKeyPair() (pub []byte, priv []byte, err error) { + pk, sk, err := mode5.Scheme().GenerateKey() + if err != nil { + return nil, nil, err + } + + pub, _ = pk.MarshalBinary() + priv, _ = sk.MarshalBinary() + return pub, priv, nil +} diff --git a/go.mod b/go.mod index 6947698ca..3b6b2678e 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,7 @@ module golang.zx2c4.com/wireguard go 1.23.1 require ( + github.com/cloudflare/circl v1.6.1 golang.org/x/crypto v0.37.0 golang.org/x/net v0.39.0 golang.org/x/sys v0.32.0 @@ -11,7 +12,6 @@ require ( ) require ( - github.com/cloudflare/circl v1.6.1 // indirect github.com/google/btree v1.1.2 // indirect golang.org/x/time v0.7.0 // indirect ) From e4017b2e00ebe41c966e6a83343aa17d0f045289 Mon Sep 17 00:00:00 2001 From: Eruel6 Date: Thu, 6 Nov 2025 14:17:29 -0300 Subject: [PATCH 33/34] feat: add tests for MLDSA Signed-off-by: Mateus Franco --- device/device_test.go | 45 +++++++++++ device/mldsa_test.go | 169 ++++++++++++++++++++++++++++++++++++++++++ device/noise_test.go | 23 +++--- device/run_tests.sh | 7 ++ 4 files changed, 230 insertions(+), 14 deletions(-) create mode 100644 device/mldsa_test.go create mode 100755 device/run_tests.sh diff --git a/device/device_test.go b/device/device_test.go index 42da6e1b0..f6041a753 100644 --- a/device/device_test.go +++ b/device/device_test.go @@ -21,6 +21,8 @@ import ( "time" "github.com/cloudflare/circl/kem/kyber/kyber1024" + "github.com/cloudflare/circl/sign/dilithium/mode5" + "golang.zx2c4.com/wireguard/conn" "golang.zx2c4.com/wireguard/conn/bindtest" @@ -193,6 +195,7 @@ func genTestPair(tb testing.TB, realSocket bool) (pair testPair) { } installMLKEMKeys(tb, &pair) + installMLDSAKeys(tb, &pair) return } @@ -508,6 +511,34 @@ func installMLKEMKeys(t testing.TB, pair *testPair) { if err := pair[1].dev.IpcSet(cfgPeer1); err != nil { t.Fatal(err) } } +func installMLDSAKeys(t testing.TB, pair *testPair) { + t.Helper() + + dil := mode5.Scheme() + pk0, sk0, err := dil.GenerateKey() + if err != nil { t.Fatal(err) } + pk1, sk1, err := dil.GenerateKey() + if err != nil { t.Fatal(err) } + + pk0b, _ := pk0.MarshalBinary() + sk0b, _ := sk0.MarshalBinary() + pk1b, _ := pk1.MarshalBinary() + sk1b, _ := sk1.MarshalBinary() + + copy(pair[0].dev.staticIdentity.mldsaPrivateKey[:], sk0b) + copy(pair[1].dev.staticIdentity.mldsaPrivateKey[:], sk1b) + + var peer0, peer1 *Peer + for _, p := range pair[0].dev.peers.keyMap { peer0 = p; break } + for _, p := range pair[1].dev.peers.keyMap { peer1 = p; break } + if peer0 == nil || peer1 == nil { + t.Fatal("não foi possível localizar peers para configurar MLDSA") + } + + copy(peer0.handshake.remoteMLDSAStatic[:], pk1b) + copy(peer1.handshake.remoteMLDSAStatic[:], pk0b) +} + func TestMLKEMKeyGeneration(t *testing.T) { pub, priv, err := GenerateQuantumKeyPair() if err != nil { t.Fatal(err) } @@ -577,6 +608,20 @@ func TestNoiseHandshakeWithMLKEM(t *testing.T) { if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { t.Fatal(err) } + dil := mode5.Scheme() + pkAS, skAS, _ := dil.GenerateKey() + pkBS, skBS, _ := dil.GenerateKey() + pkASb, _ := pkAS.MarshalBinary() + pkBSb, _ := pkBS.MarshalBinary() + skASb, _ := skAS.MarshalBinary() + skBSb, _ := skBS.MarshalBinary() + + copy(devA.staticIdentity.mldsaPrivateKey[:], skASb) + copy(devB.staticIdentity.mldsaPrivateKey[:], skBSb) + + copy(peerB.handshake.remoteMLDSAStatic[:], pkBSb) + copy(peerA.handshake.remoteMLDSAStatic[:], pkASb) + peerA.Start() peerB.Start() diff --git a/device/mldsa_test.go b/device/mldsa_test.go new file mode 100644 index 000000000..89e47d43d --- /dev/null +++ b/device/mldsa_test.go @@ -0,0 +1,169 @@ +package device + +import ( + "bytes" + "testing" + "encoding/hex" + + "github.com/cloudflare/circl/sign/dilithium/mode5" + "github.com/cloudflare/circl/kem/kyber/kyber1024" +) + +func TestGenerateMLDSAKeyPair(t *testing.T) { + pub, priv, err := GenerateMLDSAKeyPair() + if err != nil { + t.Fatalf("erro gerando MLDSA: %v", err) + } + if len(pub) != MLDSAPublicKeySize || len(priv) != MLDSAPrivateKeySize { + t.Fatalf("tamanhos inválidos: pub=%d priv=%d", len(pub), len(priv)) + } + s := mode5.Scheme() + if _, err := s.UnmarshalBinaryPublicKey(pub); err != nil { + t.Fatalf("publicKey inválida: %v", err) + } + if _, err := s.UnmarshalBinaryPrivateKey(priv); err != nil { + t.Fatalf("privateKey inválida: %v", err) + } +} + +func TestMLDSASignVerify(t *testing.T) { + s := mode5.Scheme() + pk, sk, err := s.GenerateKey() + if err != nil { + t.Fatalf("erro gerando par mldsa: %v", err) + } + msg := []byte("wireguard + mldsa test") + sig := s.Sign(sk, msg, nil) + + if !s.Verify(pk, msg, sig, nil) { + t.Fatalf("assinatura MLDSA não verificou") + } + if s.Verify(pk, append(msg, 0x01), sig, nil) { + t.Fatalf("assinatura deveria falhar em msg alterada") + } +} + +func mustCopy(dst []byte, src []byte) { + if len(dst) != len(src) { panic("tam inválido") } + copy(dst, src) +} + +func TestHybridHandshakeWithMLDSASignature(t *testing.T) { + dev1 := randDevice(t) + dev2 := randDevice(t) + defer dev1.Close() + defer dev2.Close() + + kyb := kyber1024.Scheme() + pkK1, skK1, _ := kyb.GenerateKeyPair() + pkK2, skK2, _ := kyb.GenerateKeyPair() + pubK1, _ := pkK1.MarshalBinary() + privK1, _ := skK1.MarshalBinary() + pubK2, _ := pkK2.MarshalBinary() + privK2, _ := skK2.MarshalBinary() + + mldsa := mode5.Scheme() + pkS1, skS1, _ := mldsa.GenerateKey() + pkS2, skS2, _ := mldsa.GenerateKey() + pubS1, _ := pkS1.MarshalBinary() + privS1, _ := skS1.MarshalBinary() + pubS2, _ := pkS2.MarshalBinary() + privS2, _ := skS2.MarshalBinary() + + mustCopy(dev1.staticIdentity.mlkemPrivateKey[:], privK1) + mustCopy(dev2.staticIdentity.mlkemPrivateKey[:], privK2) + mustCopy(dev1.staticIdentity.mldsaPrivateKey[:], privS1) + mustCopy(dev2.staticIdentity.mldsaPrivateKey[:], privS2) + + peer1, err := dev2.NewPeer(dev1.staticIdentity.privateKey.publicKey()) + if err != nil { t.Fatal(err) } + peer2, err := dev1.NewPeer(dev2.staticIdentity.privateKey.publicKey()) + if err != nil { t.Fatal(err) } + + mustCopy(peer1.handshake.remoteMLKEMStatic[:], pubK1) + mustCopy(peer2.handshake.remoteMLKEMStatic[:], pubK2) + mustCopy(peer1.handshake.remoteMLDSAStatic[:], pubS1) + mustCopy(peer2.handshake.remoteMLDSAStatic[:], pubS2) + + peer1.Start() + peer2.Start() + + init, err := dev1.CreateMessageInitiation(peer2) + if err != nil { + t.Fatalf("CreateMessageInitiation falhou: %v", err) + } + if p := dev2.ConsumeMessageInitiation(init); p == nil { + t.Fatalf("ConsumeMessageInitiation falhou (assinatura/MLKEM?)") + } + + resp, err := dev2.CreateMessageResponse(peer1) + if err != nil { + t.Fatalf("CreateMessageResponse falhou: %v", err) + } + if p := dev1.ConsumeMessageResponse(resp); p == nil { + t.Fatalf("ConsumeMessageResponse falhou") + } + + if err := peer1.BeginSymmetricSession(); err != nil { + t.Fatalf("peer1.BeginSymmetricSession: %v", err) + } + if err := peer2.BeginSymmetricSession(); err != nil { + t.Fatalf("peer2.BeginSymmetricSession: %v", err) + } + + key1 := peer1.keypairs.next.Load() + key2 := peer2.keypairs.current + plain := []byte("ok mldsa+mlkem+noise") + var nonce [12]byte + c := key1.send.Seal(nil, nonce[:], plain, nil) + out, err := key2.receive.Open(nil, nonce[:], c, nil) + if err != nil || !bytes.Equal(out, plain) { + t.Fatalf("falha cifrar/decifrar: %v", err) + } +} + +func TestHybridHandshake_MLDSAInvalidSignature(t *testing.T) { + dev1 := randDevice(t) + dev2 := randDevice(t) + defer dev1.Close(); defer dev2.Close() + + kyb := kyber1024.Scheme() + pkK1, skK1, _ := kyb.GenerateKeyPair() + pkK2, skK2, _ := kyb.GenerateKeyPair() + pubK1, _ := pkK1.MarshalBinary() + privK1, _ := skK1.MarshalBinary() + pubK2, _ := pkK2.MarshalBinary() + privK2, _ := skK2.MarshalBinary() + mustCopy(dev1.staticIdentity.mlkemPrivateKey[:], privK1) + mustCopy(dev2.staticIdentity.mlkemPrivateKey[:], privK2) + + s := mode5.Scheme() + pkGood, skGood, _ := s.GenerateKey() + pkWrong, _, _ := s.GenerateKey() + privGood, _ := skGood.MarshalBinary() + pubGood, _ := pkGood.MarshalBinary() + pubWrong, _ := pkWrong.MarshalBinary() + mustCopy(dev1.staticIdentity.mldsaPrivateKey[:], privGood) + + peer1, _ := dev2.NewPeer(dev1.staticIdentity.privateKey.publicKey()) + peer2, _ := dev1.NewPeer(dev2.staticIdentity.privateKey.publicKey()) + peer1.Start(); peer2.Start() + mustCopy(peer1.handshake.remoteMLKEMStatic[:], pubK1) + mustCopy(peer2.handshake.remoteMLKEMStatic[:], pubK2) + + mustCopy(peer1.handshake.remoteMLDSAStatic[:], pubWrong) + + init, err := dev1.CreateMessageInitiation(peer2) + if err != nil { + t.Fatalf("CreateMessageInitiation falhou: %v", err) + } + if p := dev2.ConsumeMessageInitiation(init); p != nil { + t.Fatalf("assinatura inválida deveria falhar") + } + + mustCopy(peer1.handshake.remoteMLDSAStatic[:], pubGood) + if p := dev2.ConsumeMessageInitiation(init); p == nil { + t.Fatalf("deveria aceitar com a pública correta") + } + _ = hex.EncodeToString +} diff --git a/device/noise_test.go b/device/noise_test.go index f0928ac66..e74e74a29 100644 --- a/device/noise_test.go +++ b/device/noise_test.go @@ -1,8 +1,3 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - package device import ( @@ -71,6 +66,15 @@ func TestNoiseHandshake(t *testing.T) { if err != nil { t.Fatal(err) } + + pair := testPair{} + pair[0].dev = dev1 + pair[1].dev = dev2 + + installMLKEMKeys(t, &pair) + + installMLDSAKeys(t, &pair) + peer1.Start() peer2.Start() @@ -80,10 +84,6 @@ func TestNoiseHandshake(t *testing.T) { peer2.handshake.precomputedStaticStatic[:], ) - /* simulate handshake */ - - // initiation message - t.Log("exchange initiation message") msg1, err := dev1.CreateMessageInitiation(peer2) @@ -110,7 +110,6 @@ func TestNoiseHandshake(t *testing.T) { peer2.handshake.hash[:], ) - // response message t.Log("exchange response message") @@ -134,8 +133,6 @@ func TestNoiseHandshake(t *testing.T) { peer2.handshake.hash[:], ) - // key pairs - t.Log("deriving keys") err = peer1.BeginSymmetricSession() @@ -151,8 +148,6 @@ func TestNoiseHandshake(t *testing.T) { key1 := peer1.keypairs.next.Load() key2 := peer2.keypairs.current - // encrypting / decryption test - t.Log("test key pairs") func() { diff --git a/device/run_tests.sh b/device/run_tests.sh new file mode 100755 index 000000000..6b81ecd2c --- /dev/null +++ b/device/run_tests.sh @@ -0,0 +1,7 @@ +cd /home/arthur/Documentos/UNB/TCC/wireguard-go/device + +TESTS=$(grep -hEo '^func[[:space:]]+Test[[:alnum:]_]*' device_test.go noise_test.go mldsa_test.go \ + | awk '{print $2}' \ + | paste -sd'|' -) + +go test -v -run "^($TESTS)$" From b2e63c1ea332e09d1f5f4a34f0a8f9bef5cef798 Mon Sep 17 00:00:00 2001 From: Eruel6 Date: Sat, 15 Nov 2025 14:36:31 -0300 Subject: [PATCH 34/34] fix: add MLDSA to benchmark Signed-off-by: Mateus Franco --- device/mlkem_bench_test.go | 235 +++++++++++++++++++++++++++---------- 1 file changed, 171 insertions(+), 64 deletions(-) diff --git a/device/mlkem_bench_test.go b/device/mlkem_bench_test.go index 349f205e9..465846221 100644 --- a/device/mlkem_bench_test.go +++ b/device/mlkem_bench_test.go @@ -94,6 +94,33 @@ func BenchmarkHandshakeWithMLKEM(b *testing.B) { b.Fatal(err) } + mldsaPubA, mldsaPrivA, err := GenerateMLDSAKeyPair() + if err != nil { + b.Fatal(err) + } + mldsaPubB, mldsaPrivB, err := GenerateMLDSAKeyPair() + if err != nil { + b.Fatal(err) + } + + devA.staticIdentity.Lock() + copy(devA.staticIdentity.mldsaPrivateKey[:], mldsaPrivA) + copy(devA.staticIdentity.mldsaPublicKey[:], mldsaPubA) + devA.staticIdentity.Unlock() + + devB.staticIdentity.Lock() + copy(devB.staticIdentity.mldsaPrivateKey[:], mldsaPrivB) + copy(devB.staticIdentity.mldsaPublicKey[:], mldsaPubB) + devB.staticIdentity.Unlock() + + peerB.handshake.mutex.Lock() + copy(peerB.handshake.remoteMLDSAStatic[:], mldsaPubB) + peerB.handshake.mutex.Unlock() + + peerA.handshake.mutex.Lock() + copy(peerA.handshake.remoteMLDSAStatic[:], mldsaPubA) + peerA.handshake.mutex.Unlock() + relaxFlood := func() { peerA.handshake.mutex.Lock() peerA.handshake.lastInitiationConsumption = time.Now().Add(-10 * time.Second) @@ -197,6 +224,33 @@ func BenchmarkHandshakeHybrid(b *testing.B) { b.Fatal("peer lookup failed (check IpcSet order)") } + mldsaPubA, mldsaPrivA, err := GenerateMLDSAKeyPair() + if err != nil { + b.Fatal(err) + } + mldsaPubB, mldsaPrivB, err := GenerateMLDSAKeyPair() + if err != nil { + b.Fatal(err) + } + + devA.staticIdentity.Lock() + copy(devA.staticIdentity.mldsaPrivateKey[:], mldsaPrivA) + copy(devA.staticIdentity.mldsaPublicKey[:], mldsaPubA) + devA.staticIdentity.Unlock() + + devB.staticIdentity.Lock() + copy(devB.staticIdentity.mldsaPrivateKey[:], mldsaPrivB) + copy(devB.staticIdentity.mldsaPublicKey[:], mldsaPubB) + devB.staticIdentity.Unlock() + + peerB.handshake.mutex.Lock() + copy(peerB.handshake.remoteMLDSAStatic[:], mldsaPubB) + peerB.handshake.mutex.Unlock() + + peerA.handshake.mutex.Lock() + copy(peerA.handshake.remoteMLDSAStatic[:], mldsaPubA) + peerA.handshake.mutex.Unlock() + relax := func() { peerA.handshake.mutex.Lock() peerA.handshake.lastInitiationConsumption = time.Now().Add(-10 * time.Second) @@ -250,78 +304,32 @@ func BenchmarkDataPlaneAEAD(b *testing.B) { devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))) devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerB.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkBb))) devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(peerA.handshake.remoteStatic[:]), "mlkem_public_key", hex.EncodeToString(pkAb))) - - msg1, _ := devA.CreateMessageInitiation(peerB) - devB.ConsumeMessageInitiation(msg1) - msg2, _ := devB.CreateMessageResponse(peerA) - devA.ConsumeMessageResponse(msg2) - peerA.BeginSymmetricSession() - peerB.BeginSymmetricSession() - - keyA := peerA.keypairs.next.Load() - keyB := peerB.keypairs.current - msg := bytes.Repeat([]byte{0x42}, 128) - var nonce [12]byte - - b.ReportAllocs() - b.SetBytes(int64(len(msg))) - b.ResetTimer() - for i := 0; i < b.N; i++ { - out := keyA.send.Seal(nil, nonce[:], msg, nil) - _, err := keyB.receive.Open(nil, nonce[:], out, nil) - if err != nil { - b.Fatal(err) - } - } -} - -func BenchmarkDataPlaneAEADHybrid(b *testing.B) { - skA, _ := newPrivateKey() - skB, _ := newPrivateKey() - tunA := tuntest.NewChannelTUN() - tunB := tuntest.NewChannelTUN() - devA := NewDevice(tunA.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) - devB := NewDevice(tunB.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) - defer devA.Close() - defer devB.Close() - - if err := devA.IpcSet(uapiCfg("private_key", hex.EncodeToString(skA[:]))); err != nil { + mldsaPubA, mldsaPrivA, err := GenerateMLDSAKeyPair() + if err != nil { b.Fatal(err) } - if err := devB.IpcSet(uapiCfg("private_key", hex.EncodeToString(skB[:]))); err != nil { + mldsaPubB, mldsaPrivB, err := GenerateMLDSAKeyPair() + if err != nil { b.Fatal(err) } - scheme := kyber1024.Scheme() - pkA, skAkem, _ := scheme.GenerateKeyPair() - pkB, skBkem, _ := scheme.GenerateKeyPair() - pkAb, _ := pkA.MarshalBinary() - pkBb, _ := pkB.MarshalBinary() - skAb, _ := skAkem.MarshalBinary() - skBb, _ := skBkem.MarshalBinary() - - if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { - b.Fatal(err) - } - if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { - b.Fatal(err) - } + devA.staticIdentity.Lock() + copy(devA.staticIdentity.mldsaPrivateKey[:], mldsaPrivA) + copy(devA.staticIdentity.mldsaPublicKey[:], mldsaPubA) + devA.staticIdentity.Unlock() - pkBNoise := skB.publicKey() - pkANoise := skA.publicKey() - if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(pkBNoise[:]), "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { - b.Fatal(err) - } - if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(pkANoise[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { - b.Fatal(err) - } + devB.staticIdentity.Lock() + copy(devB.staticIdentity.mldsaPrivateKey[:], mldsaPrivB) + copy(devB.staticIdentity.mldsaPublicKey[:], mldsaPubB) + devB.staticIdentity.Unlock() - peerB := devA.LookupPeer(pkBNoise) - peerA := devB.LookupPeer(pkANoise) - if peerA == nil || peerB == nil { - b.Fatal("peer lookup failed (check IpcSet order)") - } + peerB.handshake.mutex.Lock() + copy(peerB.handshake.remoteMLDSAStatic[:], mldsaPubB) + peerB.handshake.mutex.Unlock() + peerA.handshake.mutex.Lock() + copy(peerA.handshake.remoteMLDSAStatic[:], mldsaPubA) + peerA.handshake.mutex.Unlock() msg1, _ := devA.CreateMessageInitiation(peerB) devB.ConsumeMessageInitiation(msg1) msg2, _ := devB.CreateMessageResponse(peerA) @@ -345,3 +353,102 @@ func BenchmarkDataPlaneAEADHybrid(b *testing.B) { } } } + +func BenchmarkDataPlaneAEADHybrid(b *testing.B) { + skA, _ := newPrivateKey() + skB, _ := newPrivateKey() + tunA := tuntest.NewChannelTUN() + tunB := tuntest.NewChannelTUN() + devA := NewDevice(tunA.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + devB := NewDevice(tunB.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "")) + defer devA.Close() + defer devB.Close() + + if err := devA.IpcSet(uapiCfg("private_key", hex.EncodeToString(skA[:]))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("private_key", hex.EncodeToString(skB[:]))); err != nil { + b.Fatal(err) + } + + scheme := kyber1024.Scheme() + pkA, skAkem, _ := scheme.GenerateKeyPair() + pkB, skBkem, _ := scheme.GenerateKeyPair() + pkAb, _ := pkA.MarshalBinary() + pkBb, _ := pkB.MarshalBinary() + skAb, _ := skAkem.MarshalBinary() + skBb, _ := skBkem.MarshalBinary() + + if err := devA.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skAb))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("mlkem_private_key", hex.EncodeToString(skBb))); err != nil { + b.Fatal(err) + } + + pkBNoise := skB.publicKey() + pkANoise := skA.publicKey() + if err := devA.IpcSet(uapiCfg("public_key", hex.EncodeToString(pkBNoise[:]), "mlkem_public_key", hex.EncodeToString(pkBb))); err != nil { + b.Fatal(err) + } + if err := devB.IpcSet(uapiCfg("public_key", hex.EncodeToString(pkANoise[:]), "mlkem_public_key", hex.EncodeToString(pkAb))); err != nil { + b.Fatal(err) + } + + peerB := devA.LookupPeer(pkBNoise) + peerA := devB.LookupPeer(pkANoise) + if peerA == nil || peerB == nil { + b.Fatal("peer lookup failed (check IpcSet order)") + } + + mldsaPubA, mldsaPrivA, err := GenerateMLDSAKeyPair() + if err != nil { + b.Fatal(err) + } + mldsaPubB, mldsaPrivB, err := GenerateMLDSAKeyPair() + if err != nil { + b.Fatal(err) + } + + devA.staticIdentity.Lock() + copy(devA.staticIdentity.mldsaPrivateKey[:], mldsaPrivA) + copy(devA.staticIdentity.mldsaPublicKey[:], mldsaPubA) + devA.staticIdentity.Unlock() + + devB.staticIdentity.Lock() + copy(devB.staticIdentity.mldsaPrivateKey[:], mldsaPrivB) + copy(devB.staticIdentity.mldsaPublicKey[:], mldsaPubB) + devB.staticIdentity.Unlock() + + peerB.handshake.mutex.Lock() + copy(peerB.handshake.remoteMLDSAStatic[:], mldsaPubB) + peerB.handshake.mutex.Unlock() + + peerA.handshake.mutex.Lock() + copy(peerA.handshake.remoteMLDSAStatic[:], mldsaPubA) + peerA.handshake.mutex.Unlock() + + msg1, _ := devA.CreateMessageInitiation(peerB) + devB.ConsumeMessageInitiation(msg1) + msg2, _ := devB.CreateMessageResponse(peerA) + devA.ConsumeMessageResponse(msg2) + peerA.BeginSymmetricSession() + peerB.BeginSymmetricSession() + + keyA := peerA.keypairs.next.Load() + keyB := peerB.keypairs.current + msg := bytes.Repeat([]byte{0x42}, 128) + var nonce [12]byte + + b.ReportAllocs() + b.SetBytes(int64(len(msg))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + out := keyA.send.Seal(nil, nonce[:], msg, nil) + _, err := keyB.receive.Open(nil, nonce[:], out, nil) + if err != nil { + b.Fatal(err) + } + } +} +