From 113c8f1340ed8748da68528112d746469de11ae0 Mon Sep 17 00:00:00 2001 From: Jordan Whited Date: Thu, 27 Jun 2024 08:43:41 -0700 Subject: [PATCH 01/34] device: fix WaitPool sync.Cond usage The sync.Locker used with a sync.Cond must be acquired when changing the associated condition, otherwise there is a window within sync.Cond.Wait() where a wake-up may be missed. Fixes: 4846070 ("device: use a waiting sync.Pool instead of a channel") Reviewed-by: Brad Fitzpatrick Signed-off-by: Jordan Whited Signed-off-by: Jason A. Donenfeld --- device/pools.go | 11 ++++++----- device/pools_test.go | 4 +++- 2 files changed, 9 insertions(+), 6 deletions(-) diff --git a/device/pools.go b/device/pools.go index 94f3dc7..55d2be7 100644 --- a/device/pools.go +++ b/device/pools.go @@ -7,14 +7,13 @@ package device import ( "sync" - "sync/atomic" ) type WaitPool struct { pool sync.Pool cond sync.Cond lock sync.Mutex - count atomic.Uint32 + count uint32 // Get calls not yet Put back max uint32 } @@ -27,10 +26,10 @@ func NewWaitPool(max uint32, new func() any) *WaitPool { func (p *WaitPool) Get() any { if p.max != 0 { p.lock.Lock() - for p.count.Load() >= p.max { + for p.count >= p.max { p.cond.Wait() } - p.count.Add(1) + p.count++ p.lock.Unlock() } return p.pool.Get() @@ -41,7 +40,9 @@ func (p *WaitPool) Put(x any) { if p.max == 0 { return } - p.count.Add(^uint32(0)) + p.lock.Lock() + defer p.lock.Unlock() + p.count-- p.cond.Signal() } diff --git a/device/pools_test.go b/device/pools_test.go index 82d7493..538230b 100644 --- a/device/pools_test.go +++ b/device/pools_test.go @@ -32,7 +32,9 @@ func TestWaitPool(t *testing.T) { wg.Add(workers) var max atomic.Uint32 updateMax := func() { - count := p.count.Load() + p.lock.Lock() + count := p.count + p.lock.Unlock() if count > p.max { t.Errorf("count (%d) > max (%d)", count, p.max) } From 867a4c4a3f3a1a8fc4934553c8091c094ed6bdf8 Mon Sep 17 00:00:00 2001 From: Jordan Whited Date: Thu, 27 Jun 2024 09:06:40 -0700 Subject: [PATCH 02/34] device: fix missed return of QueueOutboundElementsContainer to its WaitPool Fixes: 3bb8fec ("conn, device, tun: implement vectorized I/O plumbing") Reviewed-by: Brad Fitzpatrick Signed-off-by: Jordan Whited Signed-off-by: Jason A. Donenfeld --- device/send.go | 1 + 1 file changed, 1 insertion(+) diff --git a/device/send.go b/device/send.go index 769720a..b20b3c5 100644 --- a/device/send.go +++ b/device/send.go @@ -506,6 +506,7 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) { device.PutMessageBuffer(elem.buffer) device.PutOutboundElement(elem) } + device.PutOutboundElementsContainer(elemsContainer) continue } dataSent := false From 9eb3221f1de589e5dd6a1721fdd7dc0fde0eb10b Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Sun, 4 May 2025 17:48:53 +0200 Subject: [PATCH 03/34] global: bump copyright notice Signed-off-by: Jason A. Donenfeld --- README.md | 2 +- conn/bind_std.go | 2 +- conn/bind_windows.go | 2 +- conn/bindtest/bindtest.go | 2 +- conn/boundif_android.go | 2 +- conn/conn.go | 2 +- conn/conn_test.go | 2 +- conn/controlfns.go | 2 +- conn/controlfns_linux.go | 2 +- conn/controlfns_unix.go | 2 +- conn/controlfns_windows.go | 2 +- conn/default.go | 2 +- conn/errors_default.go | 2 +- conn/errors_linux.go | 2 +- conn/features_default.go | 2 +- conn/features_linux.go | 2 +- conn/gso_default.go | 2 +- conn/gso_linux.go | 2 +- conn/mark_default.go | 2 +- conn/mark_unix.go | 2 +- conn/sticky_default.go | 2 +- conn/sticky_linux.go | 2 +- conn/sticky_linux_test.go | 2 +- conn/winrio/rio_windows.go | 2 +- device/allowedips.go | 2 +- device/allowedips_rand_test.go | 2 +- device/allowedips_test.go | 2 +- device/bind_test.go | 2 +- device/channels.go | 2 +- device/constants.go | 2 +- device/cookie.go | 2 +- device/cookie_test.go | 2 +- device/device.go | 2 +- device/device_test.go | 2 +- device/endpoint_test.go | 2 +- device/indextable.go | 2 +- device/ip.go | 2 +- device/kdf_test.go | 2 +- device/keypair.go | 2 +- device/logger.go | 2 +- device/mobilequirks.go | 2 +- device/noise-helpers.go | 2 +- device/noise-protocol.go | 2 +- device/noise-types.go | 2 +- device/noise_test.go | 2 +- device/peer.go | 2 +- device/pools.go | 2 +- device/pools_test.go | 2 +- device/queueconstants_android.go | 2 +- device/queueconstants_default.go | 2 +- device/queueconstants_ios.go | 2 +- device/queueconstants_windows.go | 2 +- device/race_disabled_test.go | 2 +- device/race_enabled_test.go | 2 +- device/receive.go | 2 +- device/send.go | 2 +- device/sticky_linux.go | 2 +- device/timers.go | 2 +- device/tun.go | 2 +- device/uapi.go | 2 +- format_test.go | 2 +- ipc/uapi_bsd.go | 2 +- ipc/uapi_linux.go | 2 +- ipc/uapi_unix.go | 2 +- ipc/uapi_wasm.go | 2 +- ipc/uapi_windows.go | 2 +- main.go | 2 +- main_windows.go | 2 +- ratelimiter/ratelimiter.go | 2 +- ratelimiter/ratelimiter_test.go | 2 +- replay/replay.go | 2 +- replay/replay_test.go | 2 +- rwcancel/rwcancel.go | 2 +- tai64n/tai64n.go | 2 +- tai64n/tai64n_test.go | 2 +- tun/alignment_windows_test.go | 2 +- tun/netstack/examples/http_client.go | 2 +- tun/netstack/examples/http_server.go | 2 +- tun/netstack/examples/ping_client.go | 2 +- tun/netstack/tun.go | 2 +- tun/offload_linux.go | 2 +- tun/offload_linux_test.go | 2 +- tun/operateonfd.go | 2 +- tun/tun.go | 2 +- tun/tun_darwin.go | 2 +- tun/tun_freebsd.go | 2 +- tun/tun_linux.go | 2 +- tun/tun_openbsd.go | 2 +- tun/tun_windows.go | 2 +- tun/tuntest/tuntest.go | 2 +- 90 files changed, 90 insertions(+), 90 deletions(-) diff --git a/README.md b/README.md index 074f7ec..709728d 100644 --- a/README.md +++ b/README.md @@ -56,7 +56,7 @@ $ make ## License - Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in diff --git a/conn/bind_std.go b/conn/bind_std.go index 46df7fd..f5c8816 100644 --- a/conn/bind_std.go +++ b/conn/bind_std.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/bind_windows.go b/conn/bind_windows.go index d5095e0..a3b8460 100644 --- a/conn/bind_windows.go +++ b/conn/bind_windows.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/bindtest/bindtest.go b/conn/bindtest/bindtest.go index 74e7add..46e20e6 100644 --- a/conn/bindtest/bindtest.go +++ b/conn/bindtest/bindtest.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package bindtest diff --git a/conn/boundif_android.go b/conn/boundif_android.go index dd3ca5b..be69b2a 100644 --- a/conn/boundif_android.go +++ b/conn/boundif_android.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/conn.go b/conn/conn.go index a1f57d2..1304657 100644 --- a/conn/conn.go +++ b/conn/conn.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ // Package conn implements WireGuard's network connections. diff --git a/conn/conn_test.go b/conn/conn_test.go index c6194ee..618d02b 100644 --- a/conn/conn_test.go +++ b/conn/conn_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/controlfns.go b/conn/controlfns.go index 4f7d90f..27421bd 100644 --- a/conn/controlfns.go +++ b/conn/controlfns.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/controlfns_linux.go b/conn/controlfns_linux.go index f6ab1d2..3447349 100644 --- a/conn/controlfns_linux.go +++ b/conn/controlfns_linux.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/controlfns_unix.go b/conn/controlfns_unix.go index 91692c0..b2e7570 100644 --- a/conn/controlfns_unix.go +++ b/conn/controlfns_unix.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/controlfns_windows.go b/conn/controlfns_windows.go index c3bdf7d..5e38305 100644 --- a/conn/controlfns_windows.go +++ b/conn/controlfns_windows.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/default.go b/conn/default.go index b6f761b..2ce1579 100644 --- a/conn/default.go +++ b/conn/default.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/errors_default.go b/conn/errors_default.go index f1e5b90..d967518 100644 --- a/conn/errors_default.go +++ b/conn/errors_default.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/errors_linux.go b/conn/errors_linux.go index 8e61000..037d820 100644 --- a/conn/errors_linux.go +++ b/conn/errors_linux.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/features_default.go b/conn/features_default.go index d53ff5f..cae2bea 100644 --- a/conn/features_default.go +++ b/conn/features_default.go @@ -3,7 +3,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/features_linux.go b/conn/features_linux.go index 8959d93..6386023 100644 --- a/conn/features_linux.go +++ b/conn/features_linux.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/gso_default.go b/conn/gso_default.go index 57780db..a9a3e80 100644 --- a/conn/gso_default.go +++ b/conn/gso_default.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/gso_linux.go b/conn/gso_linux.go index 8596b29..4ee31fa 100644 --- a/conn/gso_linux.go +++ b/conn/gso_linux.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/mark_default.go b/conn/mark_default.go index 3102384..72b266e 100644 --- a/conn/mark_default.go +++ b/conn/mark_default.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/mark_unix.go b/conn/mark_unix.go index d9e46ee..d0580d5 100644 --- a/conn/mark_unix.go +++ b/conn/mark_unix.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/sticky_default.go b/conn/sticky_default.go index 0b21386..15b65af 100644 --- a/conn/sticky_default.go +++ b/conn/sticky_default.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/sticky_linux.go b/conn/sticky_linux.go index 8e206e9..adfedc1 100644 --- a/conn/sticky_linux.go +++ b/conn/sticky_linux.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/sticky_linux_test.go b/conn/sticky_linux_test.go index d2bd584..1b1ee68 100644 --- a/conn/sticky_linux_test.go +++ b/conn/sticky_linux_test.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package conn diff --git a/conn/winrio/rio_windows.go b/conn/winrio/rio_windows.go index d1037bb..c396658 100644 --- a/conn/winrio/rio_windows.go +++ b/conn/winrio/rio_windows.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package winrio diff --git a/device/allowedips.go b/device/allowedips.go index fa46f97..b40c817 100644 --- a/device/allowedips.go +++ b/device/allowedips.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/allowedips_rand_test.go b/device/allowedips_rand_test.go index 07065c3..8dd9b67 100644 --- a/device/allowedips_rand_test.go +++ b/device/allowedips_rand_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/allowedips_test.go b/device/allowedips_test.go index cde068e..9ef8a76 100644 --- a/device/allowedips_test.go +++ b/device/allowedips_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/bind_test.go b/device/bind_test.go index 302a521..d3fa565 100644 --- a/device/bind_test.go +++ b/device/bind_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/channels.go b/device/channels.go index e526f6b..be15d1c 100644 --- a/device/channels.go +++ b/device/channels.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/constants.go b/device/constants.go index 59854a1..41da618 100644 --- a/device/constants.go +++ b/device/constants.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/cookie.go b/device/cookie.go index 876f05d..a093c8b 100644 --- a/device/cookie.go +++ b/device/cookie.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/cookie_test.go b/device/cookie_test.go index 4f1e50a..c937290 100644 --- a/device/cookie_test.go +++ b/device/cookie_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/device.go b/device/device.go index 83c33ee..6854ed8 100644 --- a/device/device.go +++ b/device/device.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/device_test.go b/device/device_test.go index fff172b..0091e20 100644 --- a/device/device_test.go +++ b/device/device_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/endpoint_test.go b/device/endpoint_test.go index 93a4998..85482d8 100644 --- a/device/endpoint_test.go +++ b/device/endpoint_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/indextable.go b/device/indextable.go index 00ade7d..2460fa6 100644 --- a/device/indextable.go +++ b/device/indextable.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/ip.go b/device/ip.go index eaf2363..f558744 100644 --- a/device/ip.go +++ b/device/ip.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/kdf_test.go b/device/kdf_test.go index f9c76d6..325db59 100644 --- a/device/kdf_test.go +++ b/device/kdf_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/keypair.go b/device/keypair.go index e3540d7..0b72e19 100644 --- a/device/keypair.go +++ b/device/keypair.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/logger.go b/device/logger.go index 22b0df0..a2adea3 100644 --- a/device/logger.go +++ b/device/logger.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/mobilequirks.go b/device/mobilequirks.go index 0a0080e..af4be31 100644 --- a/device/mobilequirks.go +++ b/device/mobilequirks.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/noise-helpers.go b/device/noise-helpers.go index c2f356b..35dd907 100644 --- a/device/noise-helpers.go +++ b/device/noise-helpers.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/noise-protocol.go b/device/noise-protocol.go index e8f6145..b72acf8 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/noise-types.go b/device/noise-types.go index e850359..41c944e 100644 --- a/device/noise-types.go +++ b/device/noise-types.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/noise_test.go b/device/noise_test.go index 2dd5324..f0928ac 100644 --- a/device/noise_test.go +++ b/device/noise_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/peer.go b/device/peer.go index 47a2f14..ebf25f9 100644 --- a/device/peer.go +++ b/device/peer.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/pools.go b/device/pools.go index 55d2be7..2c18f41 100644 --- a/device/pools.go +++ b/device/pools.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/pools_test.go b/device/pools_test.go index 538230b..8381d5a 100644 --- a/device/pools_test.go +++ b/device/pools_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/queueconstants_android.go b/device/queueconstants_android.go index 25f700a..236dea1 100644 --- a/device/queueconstants_android.go +++ b/device/queueconstants_android.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/queueconstants_default.go b/device/queueconstants_default.go index ea763d0..b061185 100644 --- a/device/queueconstants_default.go +++ b/device/queueconstants_default.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/queueconstants_ios.go b/device/queueconstants_ios.go index acd3cec..632e29d 100644 --- a/device/queueconstants_ios.go +++ b/device/queueconstants_ios.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/queueconstants_windows.go b/device/queueconstants_windows.go index 1eee32b..9a296d6 100644 --- a/device/queueconstants_windows.go +++ b/device/queueconstants_windows.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/race_disabled_test.go b/device/race_disabled_test.go index bb5c450..14b3284 100644 --- a/device/race_disabled_test.go +++ b/device/race_disabled_test.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/race_enabled_test.go b/device/race_enabled_test.go index 4e9daea..f1ea5cf 100644 --- a/device/race_enabled_test.go +++ b/device/race_enabled_test.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/receive.go b/device/receive.go index 1ab3e29..c7b6c87 100644 --- a/device/receive.go +++ b/device/receive.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/send.go b/device/send.go index b20b3c5..38f55c2 100644 --- a/device/send.go +++ b/device/send.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/sticky_linux.go b/device/sticky_linux.go index 6057ff1..7307b7e 100644 --- a/device/sticky_linux.go +++ b/device/sticky_linux.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. * * This implements userspace semantics of "sticky sockets", modeled after * WireGuard's kernelspace implementation. This is more or less a straight port diff --git a/device/timers.go b/device/timers.go index d4a4ed4..32519aa 100644 --- a/device/timers.go +++ b/device/timers.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. * * This is based heavily on timers.c from the kernel implementation. */ diff --git a/device/tun.go b/device/tun.go index 2a2ace9..c85dd50 100644 --- a/device/tun.go +++ b/device/tun.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/device/uapi.go b/device/uapi.go index d81dae3..521a741 100644 --- a/device/uapi.go +++ b/device/uapi.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package device diff --git a/format_test.go b/format_test.go index 6f6cab7..4d02c48 100644 --- a/format_test.go +++ b/format_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package main diff --git a/ipc/uapi_bsd.go b/ipc/uapi_bsd.go index ddcaf27..fd433a5 100644 --- a/ipc/uapi_bsd.go +++ b/ipc/uapi_bsd.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package ipc diff --git a/ipc/uapi_linux.go b/ipc/uapi_linux.go index 1562a18..fddded0 100644 --- a/ipc/uapi_linux.go +++ b/ipc/uapi_linux.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package ipc diff --git a/ipc/uapi_unix.go b/ipc/uapi_unix.go index e67be26..dcce167 100644 --- a/ipc/uapi_unix.go +++ b/ipc/uapi_unix.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package ipc diff --git a/ipc/uapi_wasm.go b/ipc/uapi_wasm.go index fa84684..50ac091 100644 --- a/ipc/uapi_wasm.go +++ b/ipc/uapi_wasm.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package ipc diff --git a/ipc/uapi_windows.go b/ipc/uapi_windows.go index aa023c9..86e60b0 100644 --- a/ipc/uapi_windows.go +++ b/ipc/uapi_windows.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package ipc diff --git a/main.go b/main.go index e016116..b6989e2 100644 --- a/main.go +++ b/main.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package main diff --git a/main_windows.go b/main_windows.go index a4dc46f..67036cf 100644 --- a/main_windows.go +++ b/main_windows.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package main diff --git a/ratelimiter/ratelimiter.go b/ratelimiter/ratelimiter.go index f7d05ef..ac69e3a 100644 --- a/ratelimiter/ratelimiter.go +++ b/ratelimiter/ratelimiter.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package ratelimiter diff --git a/ratelimiter/ratelimiter_test.go b/ratelimiter/ratelimiter_test.go index 0bfa3af..71140da 100644 --- a/ratelimiter/ratelimiter_test.go +++ b/ratelimiter/ratelimiter_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package ratelimiter diff --git a/replay/replay.go b/replay/replay.go index 8b99e23..46e224d 100644 --- a/replay/replay.go +++ b/replay/replay.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ // Package replay implements an efficient anti-replay algorithm as specified in RFC 6479. diff --git a/replay/replay_test.go b/replay/replay_test.go index 9a9e4a8..8378ec3 100644 --- a/replay/replay_test.go +++ b/replay/replay_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package replay diff --git a/rwcancel/rwcancel.go b/rwcancel/rwcancel.go index e397c0e..793e764 100644 --- a/rwcancel/rwcancel.go +++ b/rwcancel/rwcancel.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ // Package rwcancel implements cancelable read/write operations on diff --git a/tai64n/tai64n.go b/tai64n/tai64n.go index 8f10b39..e1a97a5 100644 --- a/tai64n/tai64n.go +++ b/tai64n/tai64n.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tai64n diff --git a/tai64n/tai64n_test.go b/tai64n/tai64n_test.go index c70fc1a..d0b4425 100644 --- a/tai64n/tai64n_test.go +++ b/tai64n/tai64n_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tai64n diff --git a/tun/alignment_windows_test.go b/tun/alignment_windows_test.go index 67a785e..e3252b2 100644 --- a/tun/alignment_windows_test.go +++ b/tun/alignment_windows_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/netstack/examples/http_client.go b/tun/netstack/examples/http_client.go index ccd32ed..d71267d 100644 --- a/tun/netstack/examples/http_client.go +++ b/tun/netstack/examples/http_client.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package main diff --git a/tun/netstack/examples/http_server.go b/tun/netstack/examples/http_server.go index f5b7a8f..7278851 100644 --- a/tun/netstack/examples/http_server.go +++ b/tun/netstack/examples/http_server.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package main diff --git a/tun/netstack/examples/ping_client.go b/tun/netstack/examples/ping_client.go index 2eef0fb..d1b562f 100644 --- a/tun/netstack/examples/ping_client.go +++ b/tun/netstack/examples/ping_client.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package main diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go index 2b73054..7279cd9 100644 --- a/tun/netstack/tun.go +++ b/tun/netstack/tun.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package netstack diff --git a/tun/offload_linux.go b/tun/offload_linux.go index 9ff7fea..5f0db06 100644 --- a/tun/offload_linux.go +++ b/tun/offload_linux.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/offload_linux_test.go b/tun/offload_linux_test.go index ae55c8c..d87e636 100644 --- a/tun/offload_linux_test.go +++ b/tun/offload_linux_test.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/operateonfd.go b/tun/operateonfd.go index f1beb6d..343f754 100644 --- a/tun/operateonfd.go +++ b/tun/operateonfd.go @@ -2,7 +2,7 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/tun.go b/tun/tun.go index 0ae53d0..336d642 100644 --- a/tun/tun.go +++ b/tun/tun.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/tun_darwin.go b/tun/tun_darwin.go index c9a6c0b..407b6f2 100644 --- a/tun/tun_darwin.go +++ b/tun/tun_darwin.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/tun_freebsd.go b/tun/tun_freebsd.go index 7c65fd9..4adf3a1 100644 --- a/tun/tun_freebsd.go +++ b/tun/tun_freebsd.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/tun_linux.go b/tun/tun_linux.go index bd69cb5..1461e06 100644 --- a/tun/tun_linux.go +++ b/tun/tun_linux.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/tun_openbsd.go b/tun/tun_openbsd.go index ae571b9..5aa9070 100644 --- a/tun/tun_openbsd.go +++ b/tun/tun_openbsd.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/tun_windows.go b/tun/tun_windows.go index 2af8e3e..de65fb4 100644 --- a/tun/tun_windows.go +++ b/tun/tun_windows.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tun diff --git a/tun/tuntest/tuntest.go b/tun/tuntest/tuntest.go index d07e860..9c4564f 100644 --- a/tun/tuntest/tuntest.go +++ b/tun/tuntest/tuntest.go @@ -1,6 +1,6 @@ /* SPDX-License-Identifier: MIT * - * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. */ package tuntest From 32546a15a87f253c8d03292a3be813b176a8b4e7 Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Sun, 4 May 2025 17:50:41 +0200 Subject: [PATCH 04/34] mod: bump deps Signed-off-by: Jason A. Donenfeld --- go.mod | 14 +++++++------- go.sum | 24 ++++++++++++------------ 2 files changed, 19 insertions(+), 19 deletions(-) diff --git a/go.mod b/go.mod index 919dc49..2a80e00 100644 --- a/go.mod +++ b/go.mod @@ -1,16 +1,16 @@ module golang.zx2c4.com/wireguard -go 1.20 +go 1.23.1 require ( - golang.org/x/crypto v0.13.0 - golang.org/x/net v0.15.0 - golang.org/x/sys v0.12.0 + golang.org/x/crypto v0.37.0 + golang.org/x/net v0.39.0 + golang.org/x/sys v0.32.0 golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 - gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259 + gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c ) require ( - github.com/google/btree v1.0.1 // indirect - golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 // 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 6bcecea..61875c1 100644 --- a/go.sum +++ b/go.sum @@ -1,14 +1,14 @@ -github.com/google/btree v1.0.1 h1:gK4Kx5IaGY9CD5sPJ36FHiBJ6ZXl0kilRiiCj+jdYp4= -github.com/google/btree v1.0.1/go.mod h1:xXMiIv4Fb/0kKde4SpL7qlzvu5cMJDRkFDxJfI9uaxA= -golang.org/x/crypto v0.13.0 h1:mvySKfSWJ+UKUii46M40LOvyWfN0s2U+46/jDd0e6Ck= -golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc= -golang.org/x/net v0.15.0 h1:ugBLEUaxABaB5AJqW9enI0ACdci2RUd4eP51NTBvuJ8= -golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk= -golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o= -golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 h1:vVKdlvoWBphwdxWKrFZEuM0kGgGLxUOYcY4U/2Vjg44= -golang.org/x/time v0.0.0-20220210224613-90d013bbcef8/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ= +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= +golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc= +golang.org/x/net v0.39.0 h1:ZCu7HMWDxpXpaiKdhzIfaltL9Lp31x/3fCP11bc6/fY= +golang.org/x/net v0.39.0/go.mod h1:X7NRbYVEA+ewNkCNyJ513WmMdQ3BineSwVtN2zD/d+E= +golang.org/x/sys v0.32.0 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20= +golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= +golang.org/x/time v0.7.0 h1:ntUhktv3OPE6TgYxXWv9vKvUSJyIFJlyohwbkEwPrKQ= +golang.org/x/time v0.7.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI= -gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259 h1:TbRPT0HtzFP3Cno1zZo7yPzEEnfu8EjLfl6IU9VfqkQ= -gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259/go.mod h1:AVgIgHMwK63XvmAzWG9vLQ41YnVHN0du0tEC46fI7yY= +gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c h1:m/r7OM+Y2Ty1sgBQ7Qb27VgIMBW8ZZhT4gLnUyDIhzI= +gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g= From e3c1354d27f53462801e1b86b4275699a6f9fdac Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Sun, 4 May 2025 17:54:57 +0200 Subject: [PATCH 05/34] tun/netstack: remove usage of pkt.IsNil() Since 3c75945fd ("netstack: remove PacketBuffer.IsNil()") this has been invalid. Follow the replacement pattern of that commit. The old definition inlined to the same code anyway: func (pk *PacketBuffer) IsNil() bool { return pk == nil } Signed-off-by: Jason A. Donenfeld --- tun/netstack/tun.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go index 7279cd9..04f6986 100644 --- a/tun/netstack/tun.go +++ b/tun/netstack/tun.go @@ -155,7 +155,7 @@ func (tun *netTun) Write(buf [][]byte, offset int) (int, error) { func (tun *netTun) WriteNotify() { pkt := tun.ep.Read() - if pkt.IsNil() { + if pkt == nil { return } From 45916071ba13c8f6ec44dfc46bda0449d2d54a9f Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Mon, 5 May 2025 15:09:09 +0200 Subject: [PATCH 06/34] tun/netstack: cleanup network stack at closing time Colin's commit went a step further and protected tun.incomingPacket with a lock on shutdown, but let's see if the tun.stack.Close() call actually solves that on its own. Suggested-by: kshangx Suggested-by: Colin Adler Signed-off-by: Jason A. Donenfeld --- tun/netstack/tun.go | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go index 04f6986..a7aec9e 100644 --- a/tun/netstack/tun.go +++ b/tun/netstack/tun.go @@ -43,6 +43,7 @@ type netTun struct { ep *channel.Endpoint stack *stack.Stack events chan tun.Event + notifyHandle *channel.NotificationHandle incomingPacket chan *buffer.View mtu int dnsServers []netip.Addr @@ -70,7 +71,7 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, if tcpipErr != nil { return nil, nil, fmt.Errorf("could not enable TCP SACK: %v", tcpipErr) } - dev.ep.AddNotify(dev) + dev.notifyHandle = dev.ep.AddNotify(dev) tcpipErr = dev.stack.CreateNIC(1, dev.ep) if tcpipErr != nil { return nil, nil, fmt.Errorf("CreateNIC: %v", tcpipErr) @@ -167,13 +168,14 @@ func (tun *netTun) WriteNotify() { func (tun *netTun) Close() error { tun.stack.RemoveNIC(1) + tun.stack.Close() + tun.ep.RemoveNotify(tun.notifyHandle) + tun.ep.Close() if tun.events != nil { close(tun.events) } - tun.ep.Close() - if tun.incomingPacket != nil { close(tun.incomingPacket) } From b82c016264bf9b2089da1174795a8464f311eebf Mon Sep 17 00:00:00 2001 From: Tu Dinh Ngoc Date: Thu, 20 Jun 2024 13:28:38 +0000 Subject: [PATCH 07/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 --- 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 29a8fc8..b489c56 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 c1ccff5..4ea9b8b 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 bc30fee374479c9a327285f77a06136d29884c07 Mon Sep 17 00:00:00 2001 From: ruokeqx Date: Thu, 2 Jan 2025 20:28:33 +0800 Subject: [PATCH 08/34] tun: darwin: fetch flags and mtu from if_msghdr directly Signed-off-by: ruokeqx Signed-off-by: Jason A. Donenfeld --- 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 407b6f2..341afe3 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 77b6c824a8225328aebca78eba676ea7eb606a69 Mon Sep 17 00:00:00 2001 From: Tom Holford Date: Sun, 4 May 2025 18:49:03 +0200 Subject: [PATCH 09/34] global: replaced unused function params with _ Signed-off-by: Jason A. Donenfeld --- 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 d967518..3c9b223 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 cae2bea..9fc5088 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 9ef8a76..0ce45af 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 1038256..22e1e15 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 7307b7e..f23ff02 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 0e4482a086cb6dd9bb2baac8b538d1dbf354c136 Mon Sep 17 00:00:00 2001 From: Tom Holford Date: Sun, 4 May 2025 18:49:49 +0200 Subject: [PATCH 10/34] device: use rand.NewSource instead of rand.Seed Signed-off-by: Jason A. Donenfeld --- 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 8dd9b67..b863696 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 0ce45af..7df7da5 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 436f7fdc1670df26eee958de464cf5cb0385abec Mon Sep 17 00:00:00 2001 From: Kurnia D Win Date: Wed, 7 Jun 2023 12:41:02 +0700 Subject: [PATCH 11/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 --- rwcancel/rwcancel.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rwcancel/rwcancel.go b/rwcancel/rwcancel.go index 793e764..4372453 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 9e7529c3d2d0c54f4d5384c01645a9279e4740ae Mon Sep 17 00:00:00 2001 From: Alexander Yastrebov Date: Thu, 26 Dec 2024 20:36:53 +0100 Subject: [PATCH 12/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 --- 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 b72acf8..12368ec 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 c7b6c87..1392957 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 842888ac5c93ccc5ee6344eceaadf783fcf1e243 Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Thu, 15 May 2025 16:48:14 +0200 Subject: [PATCH 13/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 --- 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 12368ec..5f713ee 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 1571e0fbae8e1d955e05dde80071bc86880d61b3 Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Thu, 15 May 2025 16:54:03 +0200 Subject: [PATCH 14/34] version: bump snapshot Signed-off-by: Jason A. Donenfeld --- version.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/version.go b/version.go index db75bb9..80f2d4b 100644 --- a/version.go +++ b/version.go @@ -1,3 +1,3 @@ package main -const Version = "0.0.20230223" +const Version = "0.0.20250515" From 256bcbd70d5b4eaae2a9f21a9889498c0f89041c Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Tue, 20 May 2025 23:03:06 +0200 Subject: [PATCH 15/34] device: add support for removing allowedips individually This pairs with the recent change in wireguard-tools. Signed-off-by: Jason A. Donenfeld --- 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 b40c817..d15373c 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 7df7da5..a4b08a3 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 521a741..cc69488 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 264889f0bbdf9250bb8389a637dd5f38389bfe0b Mon Sep 17 00:00:00 2001 From: Alexander Yastrebov Date: Sat, 17 May 2025 11:34:30 +0200 Subject: [PATCH 16/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 --- 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 5f713ee..5cf1702 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 38f55c2..ff8f7da 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 c92064f1ce35f82bf0c5a183b54e51fd5d58ad50 Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Thu, 22 May 2025 01:33:55 +0200 Subject: [PATCH 17/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 --- conn/controlfns_linux.go | 40 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/conn/controlfns_linux.go b/conn/controlfns_linux.go index 3447349..f0deefa 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 f333402bd9cbe0f3eeb02507bd14e23d7d639280 Mon Sep 17 00:00:00 2001 From: "Jason A. Donenfeld" Date: Thu, 22 May 2025 01:45:02 +0200 Subject: [PATCH 18/34] version: bump snapshot Signed-off-by: Jason A. Donenfeld --- version.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/version.go b/version.go index 80f2d4b..d5524e8 100644 --- a/version.go +++ b/version.go @@ -1,3 +1,3 @@ package main -const Version = "0.0.20250515" +const Version = "0.0.20250522" From 905a805bc6fc8741487faa7fbae498c1183740b0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Mon, 7 Nov 2022 15:29:08 +0800 Subject: [PATCH 19/34] Update .gitignore --- .gitignore | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.gitignore b/.gitignore index e460293..c7b1915 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,3 @@ wireguard-go +/.idea/ +.DS_Store \ No newline at end of file From 680e63a4f80d2139b3ada2d831a5531d09a13774 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sat, 9 Dec 2023 16:59:55 +0800 Subject: [PATCH 20/34] Add remove unused script --- remove-unused.sh | 23 +++++++++++++++++++++++ 1 file changed, 23 insertions(+) create mode 100755 remove-unused.sh diff --git a/remove-unused.sh b/remove-unused.sh new file mode 100755 index 0000000..43f3f1b --- /dev/null +++ b/remove-unused.sh @@ -0,0 +1,23 @@ +#!/usr/bin/env bash + +set -e -o pipefail + +function remove_unused() { + git rm -rf --ignore-unmatch \ + .github \ + tests \ + *_test.go \ + **/*_test.go \ + conn/bindtest \ + tun/netstack \ + tun/tuntest \ + tun/testdata \ + main*.go \ + *.md +} + +remove_unused +remove_unused + +go mod tidy +git commit -a -m "Remove unused" From e620c55272e3d2266f289e4d20d31d21a925ca28 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 17 May 2026 19:33:01 +0800 Subject: [PATCH 21/34] Add module rename script --- reformat.sh | 9 +++++++++ rename-module.sh | 37 +++++++++++++++++++++++++++++++++++++ 2 files changed, 46 insertions(+) create mode 100755 reformat.sh create mode 100755 rename-module.sh diff --git a/reformat.sh b/reformat.sh new file mode 100755 index 0000000..8d3ac05 --- /dev/null +++ b/reformat.sh @@ -0,0 +1,9 @@ +#!/usr/bin/env bash + +set -e -o pipefail + +GO_FILES=$(find . -name "*.go" | grep -v .git) + +gofumpt -l -w $GO_FILES +gofmt -l -w $GO_FILES +gci write $GO_FILES diff --git a/rename-module.sh b/rename-module.sh new file mode 100755 index 0000000..a6bd016 --- /dev/null +++ b/rename-module.sh @@ -0,0 +1,37 @@ +#!/usr/bin/env bash + +set -e -o pipefail + +OLD_MODULE_NAME="github.com/tailscale/wireguard-go" +NEW_MODULE_NAME="github.com/sagernet/wireguard-go" + +rules=$(cat <.*) + by: $NEW_MODULE_NAME\$PATH +fix: \$NEW_IMPORT +EOF +) + +sg scan --inline-rules "$rules" -U + +sed -i "s|module $OLD_MODULE_NAME|module $NEW_MODULE_NAME|" go.mod + +go mod tidy + +./reformat.sh + +git commit -m "Rename module" -a From f853bfc5c50c7fffd11eeb79cdd7a454c164d337 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 17 May 2026 20:04:34 +0800 Subject: [PATCH 22/34] Remove unused --- README.md | 77 -- conn/bind_std_test.go | 250 ------ conn/bindtest/bindtest.go | 136 ---- conn/conn_test.go | 24 - conn/sticky_linux_test.go | 266 ------- device/allowedips_rand_test.go | 141 ---- device/allowedips_test.go | 304 -------- device/bind_test.go | 56 -- device/cookie_test.go | 190 ----- device/device_test.go | 476 ------------ device/endpoint_test.go | 49 -- device/kdf_test.go | 85 --- device/noise_test.go | 179 ----- device/pools_test.go | 141 ---- device/race_disabled_test.go | 10 - device/race_enabled_test.go | 10 - format_test.go | 51 -- go.mod | 6 - go.sum | 6 - ipc/namedpipe/namedpipe_test.go | 674 ---------------- main.go | 268 ------- main_windows.go | 99 --- ratelimiter/ratelimiter_test.go | 119 --- replay/replay_test.go | 119 --- tai64n/tai64n_test.go | 40 - tests/netns.sh | 425 ----------- tun/alignment_windows_test.go | 67 -- tun/checksum_test.go | 98 --- tun/netstack/examples/http_client.go | 54 -- tun/netstack/examples/http_server.go | 51 -- tun/netstack/examples/ping_client.go | 75 -- tun/netstack/tun.go | 1057 -------------------------- tun/offload_linux_test.go | 752 ------------------ tun/tuntest/tuntest.go | 155 ---- 34 files changed, 6510 deletions(-) delete mode 100644 README.md delete mode 100644 conn/bind_std_test.go delete mode 100644 conn/bindtest/bindtest.go delete mode 100644 conn/conn_test.go delete mode 100644 conn/sticky_linux_test.go delete mode 100644 device/allowedips_rand_test.go delete mode 100644 device/allowedips_test.go delete mode 100644 device/bind_test.go delete mode 100644 device/cookie_test.go delete mode 100644 device/device_test.go delete mode 100644 device/endpoint_test.go delete mode 100644 device/kdf_test.go delete mode 100644 device/noise_test.go delete mode 100644 device/pools_test.go delete mode 100644 device/race_disabled_test.go delete mode 100644 device/race_enabled_test.go delete mode 100644 format_test.go delete mode 100644 ipc/namedpipe/namedpipe_test.go delete mode 100644 main.go delete mode 100644 main_windows.go delete mode 100644 ratelimiter/ratelimiter_test.go delete mode 100644 replay/replay_test.go delete mode 100644 tai64n/tai64n_test.go delete mode 100755 tests/netns.sh delete mode 100644 tun/alignment_windows_test.go delete mode 100644 tun/checksum_test.go delete mode 100644 tun/netstack/examples/http_client.go delete mode 100644 tun/netstack/examples/http_server.go delete mode 100644 tun/netstack/examples/ping_client.go delete mode 100644 tun/netstack/tun.go delete mode 100644 tun/offload_linux_test.go delete mode 100644 tun/tuntest/tuntest.go diff --git a/README.md b/README.md deleted file mode 100644 index 709728d..0000000 --- a/README.md +++ /dev/null @@ -1,77 +0,0 @@ -# Go Implementation of [WireGuard](https://www.wireguard.com/) - -This is an implementation of WireGuard in Go. - -## Usage - -Most Linux kernel WireGuard users are used to adding an interface with `ip link add wg0 type wireguard`. With wireguard-go, instead simply run: - -``` -$ wireguard-go wg0 -``` - -This will create an interface and fork into the background. To remove the interface, use the usual `ip link del wg0`, or if your system does not support removing interfaces directly, you may instead remove the control socket via `rm -f /var/run/wireguard/wg0.sock`, which will result in wireguard-go shutting down. - -To run wireguard-go without forking to the background, pass `-f` or `--foreground`: - -``` -$ wireguard-go -f wg0 -``` - -When an interface is running, you may use [`wg(8)`](https://git.zx2c4.com/wireguard-tools/about/src/man/wg.8) to configure it, as well as the usual `ip(8)` and `ifconfig(8)` commands. - -To run with more logging you may set the environment variable `LOG_LEVEL=debug`. - -## Platforms - -### Linux - -This will run on Linux; however you should instead use the kernel module, which is faster and better integrated into the OS. See the [installation page](https://www.wireguard.com/install/) for instructions. - -### macOS - -This runs on macOS using the utun driver. It does not yet support sticky sockets, and won't support fwmarks because of Darwin limitations. Since the utun driver cannot have arbitrary interface names, you must either use `utun[0-9]+` for an explicit interface name or `utun` to have the kernel select one for you. If you choose `utun` as the interface name, and the environment variable `WG_TUN_NAME_FILE` is defined, then the actual name of the interface chosen by the kernel is written to the file specified by that variable. - -### Windows - -This runs on Windows, but you should instead use it from the more [fully featured Windows app](https://git.zx2c4.com/wireguard-windows/about/), which uses this as a module. - -### FreeBSD - -This will run on FreeBSD. It does not yet support sticky sockets. Fwmark is mapped to `SO_USER_COOKIE`. - -### OpenBSD - -This will run on OpenBSD. It does not yet support sticky sockets. Fwmark is mapped to `SO_RTABLE`. Since the tun driver cannot have arbitrary interface names, you must either use `tun[0-9]+` for an explicit interface name or `tun` to have the program select one for you. If you choose `tun` as the interface name, and the environment variable `WG_TUN_NAME_FILE` is defined, then the actual name of the interface chosen by the kernel is written to the file specified by that variable. - -## Building - -This requires an installation of the latest version of [Go](https://go.dev/). - -``` -$ git clone https://git.zx2c4.com/wireguard-go -$ cd wireguard-go -$ make -``` - -## License - - Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - - Permission is hereby granted, free of charge, to any person obtaining a copy of - this software and associated documentation files (the "Software"), to deal in - the Software without restriction, including without limitation the rights to - use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies - of the Software, and to permit persons to whom the Software is furnished to do - so, subject to the following conditions: - - The above copyright notice and this permission notice shall be included in all - copies or substantial portions of the Software. - - THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR - IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, - FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE - AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER - LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, - OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE - SOFTWARE. diff --git a/conn/bind_std_test.go b/conn/bind_std_test.go deleted file mode 100644 index 34a3c9a..0000000 --- a/conn/bind_std_test.go +++ /dev/null @@ -1,250 +0,0 @@ -package conn - -import ( - "encoding/binary" - "net" - "testing" - - "golang.org/x/net/ipv6" -) - -func TestStdNetBindReceiveFuncAfterClose(t *testing.T) { - bind := NewStdNetBind().(*StdNetBind) - fns, _, err := bind.Open(0) - if err != nil { - t.Fatal(err) - } - bind.Close() - bufs := make([][]byte, 1) - bufs[0] = make([]byte, 1) - sizes := make([]int, 1) - eps := make([]Endpoint, 1) - for _, fn := range fns { - // The ReceiveFuncs must not access conn-related fields on StdNetBind - // unguarded. Close() nils the conn-related fields resulting in a panic - // if they violate the mutex. - fn(bufs, sizes, eps) - } -} - -func mockSetGSOSize(control *[]byte, gsoSize uint16) { - *control = (*control)[:cap(*control)] - binary.LittleEndian.PutUint16(*control, gsoSize) -} - -func Test_coalesceMessages(t *testing.T) { - cases := []struct { - name string - buffs [][]byte - wantLens []int - wantGSO []int - }{ - { - name: "one message no coalesce", - buffs: [][]byte{ - make([]byte, 1, 1), - }, - wantLens: []int{1}, - wantGSO: []int{0}, - }, - { - name: "two messages equal len coalesce", - buffs: [][]byte{ - make([]byte, 1, 2), - make([]byte, 1, 1), - }, - wantLens: []int{2}, - wantGSO: []int{1}, - }, - { - name: "two messages unequal len coalesce", - buffs: [][]byte{ - make([]byte, 2, 3), - make([]byte, 1, 1), - }, - wantLens: []int{3}, - wantGSO: []int{2}, - }, - { - name: "three messages second unequal len coalesce", - buffs: [][]byte{ - make([]byte, 2, 3), - make([]byte, 1, 1), - make([]byte, 2, 2), - }, - wantLens: []int{3, 2}, - wantGSO: []int{2, 0}, - }, - { - name: "three messages limited cap coalesce", - buffs: [][]byte{ - make([]byte, 2, 4), - make([]byte, 2, 2), - make([]byte, 2, 2), - }, - wantLens: []int{4, 2}, - wantGSO: []int{2, 0}, - }, - } - - for _, tt := range cases { - t.Run(tt.name, func(t *testing.T) { - addr := &net.UDPAddr{ - IP: net.ParseIP("127.0.0.1").To4(), - Port: 1, - } - msgs := make([]ipv6.Message, len(tt.buffs)) - for i := range msgs { - msgs[i].Buffers = make([][]byte, 1) - msgs[i].OOB = make([]byte, 0, 2) - } - got := coalesceMessages(addr, &StdNetEndpoint{AddrPort: addr.AddrPort()}, tt.buffs, msgs, mockSetGSOSize) - if got != len(tt.wantLens) { - t.Fatalf("got len %d want: %d", got, len(tt.wantLens)) - } - for i := 0; i < got; i++ { - if msgs[i].Addr != addr { - t.Errorf("msgs[%d].Addr != passed addr", i) - } - gotLen := len(msgs[i].Buffers[0]) - if gotLen != tt.wantLens[i] { - t.Errorf("len(msgs[%d].Buffers[0]) %d != %d", i, gotLen, tt.wantLens[i]) - } - gotGSO, err := mockGetGSOSize(msgs[i].OOB) - if err != nil { - t.Fatalf("msgs[%d] getGSOSize err: %v", i, err) - } - if gotGSO != tt.wantGSO[i] { - t.Errorf("msgs[%d] gsoSize %d != %d", i, gotGSO, tt.wantGSO[i]) - } - } - }) - } -} - -func mockGetGSOSize(control []byte) (int, error) { - if len(control) < 2 { - return 0, nil - } - return int(binary.LittleEndian.Uint16(control)), nil -} - -func Test_splitCoalescedMessages(t *testing.T) { - newMsg := func(n, gso int) ipv6.Message { - msg := ipv6.Message{ - Buffers: [][]byte{make([]byte, 1<<16-1)}, - N: n, - OOB: make([]byte, 2), - } - binary.LittleEndian.PutUint16(msg.OOB, uint16(gso)) - if gso > 0 { - msg.NN = 2 - } - return msg - } - - cases := []struct { - name string - msgs []ipv6.Message - firstMsgAt int - wantNumEval int - wantMsgLens []int - wantErr bool - }{ - { - name: "second last split last empty", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(3, 1), - newMsg(0, 0), - }, - firstMsgAt: 2, - wantNumEval: 3, - wantMsgLens: []int{1, 1, 1, 0}, - wantErr: false, - }, - { - name: "second last no split last empty", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(1, 0), - newMsg(0, 0), - }, - firstMsgAt: 2, - wantNumEval: 1, - wantMsgLens: []int{1, 0, 0, 0}, - wantErr: false, - }, - { - name: "second last no split last no split", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(1, 0), - newMsg(1, 0), - }, - firstMsgAt: 2, - wantNumEval: 2, - wantMsgLens: []int{1, 1, 0, 0}, - wantErr: false, - }, - { - name: "second last no split last split", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(1, 0), - newMsg(3, 1), - }, - firstMsgAt: 2, - wantNumEval: 4, - wantMsgLens: []int{1, 1, 1, 1}, - wantErr: false, - }, - { - name: "second last split last split", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(2, 1), - newMsg(2, 1), - }, - firstMsgAt: 2, - wantNumEval: 4, - wantMsgLens: []int{1, 1, 1, 1}, - wantErr: false, - }, - { - name: "second last no split last split overflow", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(1, 0), - newMsg(4, 1), - }, - firstMsgAt: 2, - wantNumEval: 4, - wantMsgLens: []int{1, 1, 1, 1}, - wantErr: true, - }, - } - - for _, tt := range cases { - t.Run(tt.name, func(t *testing.T) { - got, err := splitCoalescedMessages(tt.msgs, 2, mockGetGSOSize) - if err != nil && !tt.wantErr { - t.Fatalf("err: %v", err) - } - if got != tt.wantNumEval { - t.Fatalf("got to eval: %d want: %d", got, tt.wantNumEval) - } - for i, msg := range tt.msgs { - if msg.N != tt.wantMsgLens[i] { - t.Fatalf("msg[%d].N: %d want: %d", i, msg.N, tt.wantMsgLens[i]) - } - } - }) - } -} diff --git a/conn/bindtest/bindtest.go b/conn/bindtest/bindtest.go deleted file mode 100644 index 46e20e6..0000000 --- a/conn/bindtest/bindtest.go +++ /dev/null @@ -1,136 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package bindtest - -import ( - "fmt" - "math/rand" - "net" - "net/netip" - "os" - - "golang.zx2c4.com/wireguard/conn" -) - -type ChannelBind struct { - rx4, tx4 *chan []byte - rx6, tx6 *chan []byte - closeSignal chan bool - source4, source6 ChannelEndpoint - target4, target6 ChannelEndpoint -} - -type ChannelEndpoint uint16 - -var ( - _ conn.Bind = (*ChannelBind)(nil) - _ conn.Endpoint = (*ChannelEndpoint)(nil) -) - -func NewChannelBinds() [2]conn.Bind { - arx4 := make(chan []byte, 8192) - brx4 := make(chan []byte, 8192) - arx6 := make(chan []byte, 8192) - brx6 := make(chan []byte, 8192) - var binds [2]ChannelBind - binds[0].rx4 = &arx4 - binds[0].tx4 = &brx4 - binds[1].rx4 = &brx4 - binds[1].tx4 = &arx4 - binds[0].rx6 = &arx6 - binds[0].tx6 = &brx6 - binds[1].rx6 = &brx6 - binds[1].tx6 = &arx6 - binds[0].target4 = ChannelEndpoint(1) - binds[1].target4 = ChannelEndpoint(2) - binds[0].target6 = ChannelEndpoint(3) - binds[1].target6 = ChannelEndpoint(4) - binds[0].source4 = binds[1].target4 - binds[0].source6 = binds[1].target6 - binds[1].source4 = binds[0].target4 - binds[1].source6 = binds[0].target6 - return [2]conn.Bind{&binds[0], &binds[1]} -} - -func (c ChannelEndpoint) ClearSrc() {} - -func (c ChannelEndpoint) SrcToString() string { return "" } - -func (c ChannelEndpoint) DstToString() string { return fmt.Sprintf("127.0.0.1:%d", c) } - -func (c ChannelEndpoint) DstToBytes() []byte { return []byte{byte(c)} } - -func (c ChannelEndpoint) DstIP() netip.Addr { return netip.AddrFrom4([4]byte{127, 0, 0, 1}) } - -func (c ChannelEndpoint) SrcIP() netip.Addr { return netip.Addr{} } - -func (c *ChannelBind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err error) { - c.closeSignal = make(chan bool) - fns = append(fns, c.makeReceiveFunc(*c.rx4)) - fns = append(fns, c.makeReceiveFunc(*c.rx6)) - if rand.Uint32()&1 == 0 { - return fns, uint16(c.source4), nil - } else { - return fns, uint16(c.source6), nil - } -} - -func (c *ChannelBind) Close() error { - if c.closeSignal != nil { - select { - case <-c.closeSignal: - default: - close(c.closeSignal) - } - } - return nil -} - -func (c *ChannelBind) BatchSize() int { return 1 } - -func (c *ChannelBind) SetMark(mark uint32) error { return nil } - -func (c *ChannelBind) makeReceiveFunc(ch chan []byte) conn.ReceiveFunc { - return func(bufs [][]byte, sizes []int, eps []conn.Endpoint) (n int, err error) { - select { - case <-c.closeSignal: - return 0, net.ErrClosed - case rx := <-ch: - copied := copy(bufs[0], rx) - sizes[0] = copied - eps[0] = c.target6 - return 1, nil - } - } -} - -func (c *ChannelBind) Send(bufs [][]byte, ep conn.Endpoint) error { - for _, b := range bufs { - select { - case <-c.closeSignal: - return net.ErrClosed - default: - bc := make([]byte, len(b)) - copy(bc, b) - if ep.(ChannelEndpoint) == c.target4 { - *c.tx4 <- bc - } else if ep.(ChannelEndpoint) == c.target6 { - *c.tx6 <- bc - } else { - return os.ErrInvalid - } - } - } - return nil -} - -func (c *ChannelBind) ParseEndpoint(s string) (conn.Endpoint, error) { - addr, err := netip.ParseAddrPort(s) - if err != nil { - return nil, err - } - return ChannelEndpoint(addr.Port()), nil -} diff --git a/conn/conn_test.go b/conn/conn_test.go deleted file mode 100644 index 618d02b..0000000 --- a/conn/conn_test.go +++ /dev/null @@ -1,24 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package conn - -import ( - "testing" -) - -func TestPrettyName(t *testing.T) { - var ( - recvFunc ReceiveFunc = func(bufs [][]byte, sizes []int, eps []Endpoint) (n int, err error) { return } - ) - - const want = "TestPrettyName" - - t.Run("ReceiveFunc.PrettyName", func(t *testing.T) { - if got := recvFunc.PrettyName(); got != want { - t.Errorf("PrettyName() = %v, want %v", got, want) - } - }) -} diff --git a/conn/sticky_linux_test.go b/conn/sticky_linux_test.go deleted file mode 100644 index 1b1ee68..0000000 --- a/conn/sticky_linux_test.go +++ /dev/null @@ -1,266 +0,0 @@ -//go:build linux && !android - -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package conn - -import ( - "context" - "net" - "net/netip" - "runtime" - "testing" - "unsafe" - - "golang.org/x/sys/unix" -) - -func setSrc(ep *StdNetEndpoint, addr netip.Addr, ifidx int32) { - var buf []byte - if addr.Is4() { - buf = make([]byte, unix.CmsgSpace(unix.SizeofInet4Pktinfo)) - hdr := unix.Cmsghdr{ - Level: unix.IPPROTO_IP, - Type: unix.IP_PKTINFO, - } - hdr.SetLen(unix.CmsgLen(unix.SizeofInet4Pktinfo)) - copy(buf, unsafe.Slice((*byte)(unsafe.Pointer(&hdr)), int(unsafe.Sizeof(hdr)))) - - info := unix.Inet4Pktinfo{ - Ifindex: ifidx, - Spec_dst: addr.As4(), - } - copy(buf[unix.CmsgLen(0):], unsafe.Slice((*byte)(unsafe.Pointer(&info)), unix.SizeofInet4Pktinfo)) - } else { - buf = make([]byte, unix.CmsgSpace(unix.SizeofInet6Pktinfo)) - hdr := unix.Cmsghdr{ - Level: unix.IPPROTO_IPV6, - Type: unix.IPV6_PKTINFO, - } - hdr.SetLen(unix.CmsgLen(unix.SizeofInet6Pktinfo)) - copy(buf, unsafe.Slice((*byte)(unsafe.Pointer(&hdr)), int(unsafe.Sizeof(hdr)))) - - info := unix.Inet6Pktinfo{ - Ifindex: uint32(ifidx), - Addr: addr.As16(), - } - copy(buf[unix.CmsgLen(0):], unsafe.Slice((*byte)(unsafe.Pointer(&info)), unix.SizeofInet6Pktinfo)) - } - - ep.src = buf -} - -func Test_setSrcControl(t *testing.T) { - t.Run("IPv4", func(t *testing.T) { - ep := &StdNetEndpoint{ - AddrPort: netip.MustParseAddrPort("127.0.0.1:1234"), - } - setSrc(ep, netip.MustParseAddr("127.0.0.1"), 5) - - control := make([]byte, stickyControlSize) - - setSrcControl(&control, ep) - - hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0])) - if hdr.Level != unix.IPPROTO_IP { - t.Errorf("unexpected level: %d", hdr.Level) - } - if hdr.Type != unix.IP_PKTINFO { - t.Errorf("unexpected type: %d", hdr.Type) - } - if uint(hdr.Len) != uint(unix.CmsgLen(int(unsafe.Sizeof(unix.Inet4Pktinfo{})))) { - t.Errorf("unexpected length: %d", hdr.Len) - } - info := (*unix.Inet4Pktinfo)(unsafe.Pointer(&control[unix.CmsgLen(0)])) - if info.Spec_dst[0] != 127 || info.Spec_dst[1] != 0 || info.Spec_dst[2] != 0 || info.Spec_dst[3] != 1 { - t.Errorf("unexpected address: %v", info.Spec_dst) - } - if info.Ifindex != 5 { - t.Errorf("unexpected ifindex: %d", info.Ifindex) - } - }) - - t.Run("IPv6", func(t *testing.T) { - ep := &StdNetEndpoint{ - AddrPort: netip.MustParseAddrPort("[::1]:1234"), - } - setSrc(ep, netip.MustParseAddr("::1"), 5) - - control := make([]byte, stickyControlSize) - - setSrcControl(&control, ep) - - hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0])) - if hdr.Level != unix.IPPROTO_IPV6 { - t.Errorf("unexpected level: %d", hdr.Level) - } - if hdr.Type != unix.IPV6_PKTINFO { - t.Errorf("unexpected type: %d", hdr.Type) - } - if uint(hdr.Len) != uint(unix.CmsgLen(int(unsafe.Sizeof(unix.Inet6Pktinfo{})))) { - t.Errorf("unexpected length: %d", hdr.Len) - } - info := (*unix.Inet6Pktinfo)(unsafe.Pointer(&control[unix.CmsgLen(0)])) - if info.Addr != ep.SrcIP().As16() { - t.Errorf("unexpected address: %v", info.Addr) - } - if info.Ifindex != 5 { - t.Errorf("unexpected ifindex: %d", info.Ifindex) - } - }) - - t.Run("ClearOnNoSrc", func(t *testing.T) { - control := make([]byte, stickyControlSize) - hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0])) - hdr.Level = 1 - hdr.Type = 2 - hdr.Len = 3 - - setSrcControl(&control, &StdNetEndpoint{}) - - if len(control) != 0 { - t.Errorf("unexpected control: %v", control) - } - }) -} - -func Test_getSrcFromControl(t *testing.T) { - t.Run("IPv4", func(t *testing.T) { - control := make([]byte, stickyControlSize) - hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0])) - hdr.Level = unix.IPPROTO_IP - hdr.Type = unix.IP_PKTINFO - hdr.SetLen(unix.CmsgLen(int(unsafe.Sizeof(unix.Inet4Pktinfo{})))) - info := (*unix.Inet4Pktinfo)(unsafe.Pointer(&control[unix.CmsgLen(0)])) - info.Spec_dst = [4]byte{127, 0, 0, 1} - info.Ifindex = 5 - - ep := &StdNetEndpoint{} - getSrcFromControl(control, ep) - - if ep.SrcIP() != netip.MustParseAddr("127.0.0.1") { - t.Errorf("unexpected address: %v", ep.SrcIP()) - } - if ep.SrcIfidx() != 5 { - t.Errorf("unexpected ifindex: %d", ep.SrcIfidx()) - } - }) - t.Run("IPv6", func(t *testing.T) { - control := make([]byte, stickyControlSize) - hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0])) - hdr.Level = unix.IPPROTO_IPV6 - hdr.Type = unix.IPV6_PKTINFO - hdr.SetLen(unix.CmsgLen(int(unsafe.Sizeof(unix.Inet6Pktinfo{})))) - info := (*unix.Inet6Pktinfo)(unsafe.Pointer(&control[unix.CmsgLen(0)])) - info.Addr = [16]byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1} - info.Ifindex = 5 - - ep := &StdNetEndpoint{} - getSrcFromControl(control, ep) - - if ep.SrcIP() != netip.MustParseAddr("::1") { - t.Errorf("unexpected address: %v", ep.SrcIP()) - } - if ep.SrcIfidx() != 5 { - t.Errorf("unexpected ifindex: %d", ep.SrcIfidx()) - } - }) - t.Run("ClearOnEmpty", func(t *testing.T) { - var control []byte - ep := &StdNetEndpoint{} - setSrc(ep, netip.MustParseAddr("::1"), 5) - - getSrcFromControl(control, ep) - if ep.SrcIP().IsValid() { - t.Errorf("unexpected address: %v", ep.SrcIP()) - } - if ep.SrcIfidx() != 0 { - t.Errorf("unexpected ifindex: %d", ep.SrcIfidx()) - } - }) - t.Run("Multiple", func(t *testing.T) { - zeroControl := make([]byte, unix.CmsgSpace(0)) - zeroHdr := (*unix.Cmsghdr)(unsafe.Pointer(&zeroControl[0])) - zeroHdr.SetLen(unix.CmsgLen(0)) - - control := make([]byte, unix.CmsgSpace(unix.SizeofInet4Pktinfo)) - hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0])) - hdr.Level = unix.IPPROTO_IP - hdr.Type = unix.IP_PKTINFO - hdr.SetLen(unix.CmsgLen(int(unsafe.Sizeof(unix.Inet4Pktinfo{})))) - info := (*unix.Inet4Pktinfo)(unsafe.Pointer(&control[unix.CmsgLen(0)])) - info.Spec_dst = [4]byte{127, 0, 0, 1} - info.Ifindex = 5 - - combined := make([]byte, 0) - combined = append(combined, zeroControl...) - combined = append(combined, control...) - - ep := &StdNetEndpoint{} - getSrcFromControl(combined, ep) - - if ep.SrcIP() != netip.MustParseAddr("127.0.0.1") { - t.Errorf("unexpected address: %v", ep.SrcIP()) - } - if ep.SrcIfidx() != 5 { - t.Errorf("unexpected ifindex: %d", ep.SrcIfidx()) - } - }) -} - -func Test_listenConfig(t *testing.T) { - t.Run("IPv4", func(t *testing.T) { - conn, err := listenConfig().ListenPacket(context.Background(), "udp4", ":0") - if err != nil { - t.Fatal(err) - } - defer conn.Close() - sc, err := conn.(*net.UDPConn).SyscallConn() - if err != nil { - t.Fatal(err) - } - - if runtime.GOOS == "linux" { - var i int - sc.Control(func(fd uintptr) { - i, err = unix.GetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_PKTINFO) - }) - if err != nil { - t.Fatal(err) - } - if i != 1 { - t.Error("IP_PKTINFO not set!") - } - } else { - t.Logf("listenConfig() does not set IPV6_RECVPKTINFO on %s", runtime.GOOS) - } - }) - t.Run("IPv6", func(t *testing.T) { - conn, err := listenConfig().ListenPacket(context.Background(), "udp6", ":0") - if err != nil { - t.Fatal(err) - } - sc, err := conn.(*net.UDPConn).SyscallConn() - if err != nil { - t.Fatal(err) - } - - if runtime.GOOS == "linux" { - var i int - sc.Control(func(fd uintptr) { - i, err = unix.GetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_RECVPKTINFO) - }) - if err != nil { - t.Fatal(err) - } - if i != 1 { - t.Error("IPV6_PKTINFO not set!") - } - } else { - t.Logf("listenConfig() does not set IPV6_RECVPKTINFO on %s", runtime.GOOS) - } - }) -} diff --git a/device/allowedips_rand_test.go b/device/allowedips_rand_test.go deleted file mode 100644 index b863696..0000000 --- a/device/allowedips_rand_test.go +++ /dev/null @@ -1,141 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "math/rand" - "net" - "net/netip" - "sort" - "testing" -) - -const ( - NumberOfPeers = 100 - NumberOfPeerRemovals = 4 - NumberOfAddresses = 250 - NumberOfTests = 10000 -) - -type SlowNode struct { - peer *Peer - cidr uint8 - bits []byte -} - -type SlowRouter []*SlowNode - -func (r SlowRouter) Len() int { - return len(r) -} - -func (r SlowRouter) Less(i, j int) bool { - return r[i].cidr > r[j].cidr -} - -func (r SlowRouter) Swap(i, j int) { - r[i], r[j] = r[j], r[i] -} - -func (r SlowRouter) Insert(addr []byte, cidr uint8, peer *Peer) SlowRouter { - for _, t := range r { - if t.cidr == cidr && commonBits(t.bits, addr) >= cidr { - t.peer = peer - t.bits = addr - return r - } - } - r = append(r, &SlowNode{ - cidr: cidr, - bits: addr, - peer: peer, - }) - sort.Sort(r) - return r -} - -func (r SlowRouter) Lookup(addr []byte) *Peer { - for _, t := range r { - common := commonBits(t.bits, addr) - if common >= t.cidr { - return t.peer - } - } - return nil -} - -func (r SlowRouter) RemoveByPeer(peer *Peer) SlowRouter { - n := 0 - for _, x := range r { - if x.peer != peer { - r[n] = x - n++ - } - } - return r[:n] -} - -func TestTrieRandom(t *testing.T) { - var slow4, slow6 SlowRouter - var peers []*Peer - var allowedIPs AllowedIPs - - rng := rand.New(rand.NewSource(1)) - - for n := 0; n < NumberOfPeers; n++ { - peers = append(peers, &Peer{}) - } - - for n := 0; n < NumberOfAddresses; n++ { - var addr4 [4]byte - 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 - rng.Read(addr6[:]) - cidr = uint8(rand.Intn(128) + 1) - index = rand.Intn(NumberOfPeers) - allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom16(addr6), int(cidr)), peers[index]) - slow6 = slow6.Insert(addr6[:], cidr, peers[index]) - } - - var p int - for p = 0; ; p++ { - for n := 0; n < NumberOfTests; n++ { - var addr4 [4]byte - rng.Read(addr4[:]) - peer1 := slow4.Lookup(addr4[:]) - peer2 := allowedIPs.Lookup(addr4[:]) - if peer1 != peer2 { - t.Errorf("Trie did not match naive implementation, for %v: want %p, got %p", net.IP(addr4[:]), peer1, peer2) - } - - var addr6 [16]byte - rng.Read(addr6[:]) - peer1 = slow6.Lookup(addr6[:]) - peer2 = allowedIPs.Lookup(addr6[:]) - if peer1 != peer2 { - t.Errorf("Trie did not match naive implementation, for %v: want %p, got %p", net.IP(addr6[:]), peer1, peer2) - } - } - if p >= len(peers) || p >= NumberOfPeerRemovals { - break - } - allowedIPs.RemoveByPeer(peers[p]) - slow4 = slow4.RemoveByPeer(peers[p]) - slow6 = slow6.RemoveByPeer(peers[p]) - } - for ; p < len(peers); p++ { - allowedIPs.RemoveByPeer(peers[p]) - } - - if allowedIPs.IPv4 != nil || allowedIPs.IPv6 != nil { - t.Error("Failed to remove all nodes from trie by peer") - } -} diff --git a/device/allowedips_test.go b/device/allowedips_test.go deleted file mode 100644 index a4b08a3..0000000 --- a/device/allowedips_test.go +++ /dev/null @@ -1,304 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "math/rand" - "net" - "net/netip" - "testing" -) - -type testPairCommonBits struct { - s1 []byte - s2 []byte - match uint8 -} - -func TestCommonBits(t *testing.T) { - tests := []testPairCommonBits{ - {s1: []byte{1, 4, 53, 128}, s2: []byte{0, 0, 0, 0}, match: 7}, - {s1: []byte{0, 4, 53, 128}, s2: []byte{0, 0, 0, 0}, match: 13}, - {s1: []byte{0, 4, 53, 253}, s2: []byte{0, 4, 53, 252}, match: 31}, - {s1: []byte{192, 168, 1, 1}, s2: []byte{192, 169, 1, 1}, match: 15}, - {s1: []byte{65, 168, 1, 1}, s2: []byte{192, 169, 1, 1}, match: 0}, - } - - for _, p := range tests { - v := commonBits(p.s1, p.s2) - if v != p.match { - t.Error( - "For slice", p.s1, p.s2, - "expected match", p.match, - ",but got", v, - ) - } - } -} - -func benchmarkTrie(peerNumber, addressNumber, _ int, b *testing.B) { - var trie *trieEntry - var peers []*Peer - root := parentIndirection{&trie, 2} - - rng := rand.New(rand.NewSource(1)) - - const AddressLength = 4 - - for n := 0; n < peerNumber; n++ { - peers = append(peers, &Peer{}) - } - - for n := 0; n < addressNumber; n++ { - var addr [AddressLength]byte - 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 - rng.Read(addr[:]) - trie.lookup(addr[:]) - } -} - -func BenchmarkTrieIPv4Peers100Addresses1000(b *testing.B) { - benchmarkTrie(100, 1000, net.IPv4len, b) -} - -func BenchmarkTrieIPv4Peers10Addresses10(b *testing.B) { - benchmarkTrie(10, 10, net.IPv4len, b) -} - -func BenchmarkTrieIPv6Peers100Addresses1000(b *testing.B) { - benchmarkTrie(100, 1000, net.IPv6len, b) -} - -func BenchmarkTrieIPv6Peers10Addresses10(b *testing.B) { - benchmarkTrie(10, 10, net.IPv6len, b) -} - -/* Test ported from kernel implementation: - * selftest/allowedips.h - */ -func TestTrieIPv4(t *testing.T) { - a := &Peer{} - b := &Peer{} - c := &Peer{} - d := &Peer{} - e := &Peer{} - g := &Peer{} - h := &Peer{} - - var allowedIPs AllowedIPs - - insert := func(peer *Peer, a, b, c, d byte, cidr uint8) { - 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 { - t.Error("Assert EQ failed") - } - } - - assertNEQ := func(peer *Peer, a, b, c, d byte) { - p := allowedIPs.Lookup([]byte{a, b, c, d}) - if p == peer { - t.Error("Assert NEQ failed") - } - } - - insert(a, 192, 168, 4, 0, 24) - insert(b, 192, 168, 4, 4, 32) - insert(c, 192, 168, 0, 0, 16) - insert(d, 192, 95, 5, 64, 27) - insert(c, 192, 95, 5, 65, 27) - insert(e, 0, 0, 0, 0, 0) - insert(g, 64, 15, 112, 0, 20) - insert(h, 64, 15, 123, 211, 25) - insert(a, 10, 0, 0, 0, 25) - insert(b, 10, 0, 0, 128, 25) - insert(a, 10, 1, 0, 0, 30) - insert(b, 10, 1, 0, 4, 30) - insert(c, 10, 1, 0, 8, 29) - insert(d, 10, 1, 0, 16, 29) - - assertEQ(a, 192, 168, 4, 20) - assertEQ(a, 192, 168, 4, 0) - assertEQ(b, 192, 168, 4, 4) - assertEQ(c, 192, 168, 200, 182) - assertEQ(c, 192, 95, 5, 68) - assertEQ(e, 192, 95, 5, 96) - assertEQ(g, 64, 15, 116, 26) - assertEQ(g, 64, 15, 127, 3) - - insert(a, 1, 0, 0, 0, 32) - insert(a, 64, 0, 0, 0, 32) - insert(a, 128, 0, 0, 0, 32) - insert(a, 192, 0, 0, 0, 32) - insert(a, 255, 0, 0, 0, 32) - - assertEQ(a, 1, 0, 0, 0) - assertEQ(a, 64, 0, 0, 0) - assertEQ(a, 128, 0, 0, 0) - assertEQ(a, 192, 0, 0, 0) - assertEQ(a, 255, 0, 0, 0) - - allowedIPs.RemoveByPeer(a) - - assertNEQ(a, 1, 0, 0, 0) - assertNEQ(a, 64, 0, 0, 0) - assertNEQ(a, 128, 0, 0, 0) - assertNEQ(a, 192, 0, 0, 0) - assertNEQ(a, 255, 0, 0, 0) - - allowedIPs.RemoveByPeer(a) - allowedIPs.RemoveByPeer(b) - allowedIPs.RemoveByPeer(c) - allowedIPs.RemoveByPeer(d) - allowedIPs.RemoveByPeer(e) - allowedIPs.RemoveByPeer(g) - allowedIPs.RemoveByPeer(h) - if allowedIPs.IPv4 != nil || allowedIPs.IPv6 != nil { - t.Error("Expected removing all the peers to empty trie, but it did not") - } - - insert(a, 192, 168, 0, 0, 16) - insert(a, 192, 168, 0, 0, 24) - - 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: - * selftest/allowedips.h - */ -func TestTrieIPv6(t *testing.T) { - a := &Peer{} - b := &Peer{} - c := &Peer{} - d := &Peer{} - e := &Peer{} - f := &Peer{} - g := &Peer{} - h := &Peer{} - - var allowedIPs AllowedIPs - - expand := func(a uint32) []byte { - var out [4]byte - out[0] = byte(a >> 24 & 0xff) - out[1] = byte(a >> 16 & 0xff) - out[2] = byte(a >> 8 & 0xff) - out[3] = byte(a & 0xff) - return out[:] - } - - insert := 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.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)...) - 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 EQ failed") - } - } - - 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) - insert(f, 0, 0, 0, 0, 0) - insert(g, 0x24046800, 0, 0, 0, 32) - insert(h, 0x24046800, 0x40040800, 0xdeadbeef, 0xdeadbeef, 64) - insert(a, 0x24046800, 0x40040800, 0xdeadbeef, 0xdeadbeef, 128) - insert(c, 0x24446800, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128) - insert(b, 0x24446800, 0xf0e40800, 0xeeaebeef, 0, 98) - - assertEQ(d, 0x26075300, 0x60006b00, 0, 0xc05f0543) - assertEQ(c, 0x26075300, 0x60006b00, 0, 0xc02e01ee) - assertEQ(f, 0x26075300, 0x60006b01, 0, 0) - assertEQ(g, 0x24046800, 0x40040806, 0, 0x1006) - assertEQ(g, 0x24046800, 0x40040806, 0x1234, 0x5678) - assertEQ(f, 0x240467ff, 0x40040806, 0x1234, 0x5678) - assertEQ(f, 0x24046801, 0x40040806, 0x1234, 0x5678) - assertEQ(h, 0x24046800, 0x40040800, 0x1234, 0x5678) - 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/bind_test.go b/device/bind_test.go deleted file mode 100644 index d3fa565..0000000 --- a/device/bind_test.go +++ /dev/null @@ -1,56 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "errors" - - "golang.zx2c4.com/wireguard/conn" -) - -type DummyDatagram struct { - msg []byte - endpoint conn.Endpoint -} - -type DummyBind struct { - in6 chan DummyDatagram - in4 chan DummyDatagram - closed bool -} - -func (b *DummyBind) SetMark(v uint32) error { - return nil -} - -func (b *DummyBind) ReceiveIPv6(buf []byte) (int, conn.Endpoint, error) { - datagram, ok := <-b.in6 - if !ok { - return 0, nil, errors.New("closed") - } - copy(buf, datagram.msg) - return len(datagram.msg), datagram.endpoint, nil -} - -func (b *DummyBind) ReceiveIPv4(buf []byte) (int, conn.Endpoint, error) { - datagram, ok := <-b.in4 - if !ok { - return 0, nil, errors.New("closed") - } - copy(buf, datagram.msg) - return len(datagram.msg), datagram.endpoint, nil -} - -func (b *DummyBind) Close() error { - close(b.in6) - close(b.in4) - b.closed = true - return nil -} - -func (b *DummyBind) Send(buf []byte, end conn.Endpoint) error { - return nil -} diff --git a/device/cookie_test.go b/device/cookie_test.go deleted file mode 100644 index c937290..0000000 --- a/device/cookie_test.go +++ /dev/null @@ -1,190 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "testing" -) - -func TestCookieMAC1(t *testing.T) { - // setup generator / checker - - var ( - generator CookieGenerator - checker CookieChecker - ) - - sk, err := newPrivateKey() - if err != nil { - t.Fatal(err) - } - pk := sk.publicKey() - - generator.Init(pk) - checker.Init(pk) - - // check mac1 - - src := []byte{192, 168, 13, 37, 10, 10, 10} - - checkMAC1 := func(msg []byte) { - generator.AddMacs(msg) - if !checker.CheckMAC1(msg) { - t.Fatal("MAC1 generation/verification failed") - } - if checker.CheckMAC2(msg, src) { - t.Fatal("MAC2 generation/verification failed") - } - } - - checkMAC1([]byte{ - 0x99, 0xbb, 0xa5, 0xfc, 0x99, 0xaa, 0x83, 0xbd, - 0x7b, 0x00, 0xc5, 0x9a, 0x4c, 0xb9, 0xcf, 0x62, - 0x40, 0x23, 0xf3, 0x8e, 0xd8, 0xd0, 0x62, 0x64, - 0x5d, 0xb2, 0x80, 0x13, 0xda, 0xce, 0xc6, 0x91, - 0x61, 0xd6, 0x30, 0xf1, 0x32, 0xb3, 0xa2, 0xf4, - 0x7b, 0x43, 0xb5, 0xa7, 0xe2, 0xb1, 0xf5, 0x6c, - 0x74, 0x6b, 0xb0, 0xcd, 0x1f, 0x94, 0x86, 0x7b, - 0xc8, 0xfb, 0x92, 0xed, 0x54, 0x9b, 0x44, 0xf5, - 0xc8, 0x7d, 0xb7, 0x8e, 0xff, 0x49, 0xc4, 0xe8, - 0x39, 0x7c, 0x19, 0xe0, 0x60, 0x19, 0x51, 0xf8, - 0xe4, 0x8e, 0x02, 0xf1, 0x7f, 0x1d, 0xcc, 0x8e, - 0xb0, 0x07, 0xff, 0xf8, 0xaf, 0x7f, 0x66, 0x82, - 0x83, 0xcc, 0x7c, 0xfa, 0x80, 0xdb, 0x81, 0x53, - 0xad, 0xf7, 0xd8, 0x0c, 0x10, 0xe0, 0x20, 0xfd, - 0xe8, 0x0b, 0x3f, 0x90, 0x15, 0xcd, 0x93, 0xad, - 0x0b, 0xd5, 0x0c, 0xcc, 0x88, 0x56, 0xe4, 0x3f, - }) - - checkMAC1([]byte{ - 0x33, 0xe7, 0x2a, 0x84, 0x9f, 0xff, 0x57, 0x6c, - 0x2d, 0xc3, 0x2d, 0xe1, 0xf5, 0x5c, 0x97, 0x56, - 0xb8, 0x93, 0xc2, 0x7d, 0xd4, 0x41, 0xdd, 0x7a, - 0x4a, 0x59, 0x3b, 0x50, 0xdd, 0x7a, 0x7a, 0x8c, - 0x9b, 0x96, 0xaf, 0x55, 0x3c, 0xeb, 0x6d, 0x0b, - 0x13, 0x0b, 0x97, 0x98, 0xb3, 0x40, 0xc3, 0xcc, - 0xb8, 0x57, 0x33, 0x45, 0x6e, 0x8b, 0x09, 0x2b, - 0x81, 0x2e, 0xd2, 0xb9, 0x66, 0x0b, 0x93, 0x05, - }) - - checkMAC1([]byte{ - 0x9b, 0x96, 0xaf, 0x55, 0x3c, 0xeb, 0x6d, 0x0b, - 0x13, 0x0b, 0x97, 0x98, 0xb3, 0x40, 0xc3, 0xcc, - 0xb8, 0x57, 0x33, 0x45, 0x6e, 0x8b, 0x09, 0x2b, - 0x81, 0x2e, 0xd2, 0xb9, 0x66, 0x0b, 0x93, 0x05, - }) - - // exchange cookie reply - - func() { - msg := []byte{ - 0x6d, 0xd7, 0xc3, 0x2e, 0xb0, 0x76, 0xd8, 0xdf, - 0x30, 0x65, 0x7d, 0x62, 0x3e, 0xf8, 0x9a, 0xe8, - 0xe7, 0x3c, 0x64, 0xa3, 0x78, 0x48, 0xda, 0xf5, - 0x25, 0x61, 0x28, 0x53, 0x79, 0x32, 0x86, 0x9f, - 0xa0, 0x27, 0x95, 0x69, 0xb6, 0xba, 0xd0, 0xa2, - 0xf8, 0x68, 0xea, 0xa8, 0x62, 0xf2, 0xfd, 0x1b, - 0xe0, 0xb4, 0x80, 0xe5, 0x6b, 0x3a, 0x16, 0x9e, - 0x35, 0xf6, 0xa8, 0xf2, 0x4f, 0x9a, 0x7b, 0xe9, - 0x77, 0x0b, 0xc2, 0xb4, 0xed, 0xba, 0xf9, 0x22, - 0xc3, 0x03, 0x97, 0x42, 0x9f, 0x79, 0x74, 0x27, - 0xfe, 0xf9, 0x06, 0x6e, 0x97, 0x3a, 0xa6, 0x8f, - 0xc9, 0x57, 0x0a, 0x54, 0x4c, 0x64, 0x4a, 0xe2, - 0x4f, 0xa1, 0xce, 0x95, 0x9b, 0x23, 0xa9, 0x2b, - 0x85, 0x93, 0x42, 0xb0, 0xa5, 0x53, 0xed, 0xeb, - 0x63, 0x2a, 0xf1, 0x6d, 0x46, 0xcb, 0x2f, 0x61, - 0x8c, 0xe1, 0xe8, 0xfa, 0x67, 0x20, 0x80, 0x6d, - } - generator.AddMacs(msg) - reply, err := checker.CreateReply(msg, 1377, src) - if err != nil { - t.Fatal("Failed to create cookie reply:", err) - } - if !generator.ConsumeReply(reply) { - t.Fatal("Failed to consume cookie reply") - } - }() - - // check mac2 - - checkMAC2 := func(msg []byte) { - generator.AddMacs(msg) - - if !checker.CheckMAC1(msg) { - t.Fatal("MAC1 generation/verification failed") - } - if !checker.CheckMAC2(msg, src) { - t.Fatal("MAC2 generation/verification failed") - } - - msg[5] ^= 0x20 - - if checker.CheckMAC1(msg) { - t.Fatal("MAC1 generation/verification failed") - } - if checker.CheckMAC2(msg, src) { - t.Fatal("MAC2 generation/verification failed") - } - - msg[5] ^= 0x20 - - srcBad1 := []byte{192, 168, 13, 37, 40, 1} - if checker.CheckMAC2(msg, srcBad1) { - t.Fatal("MAC2 generation/verification failed") - } - - srcBad2 := []byte{192, 168, 13, 38, 40, 1} - if checker.CheckMAC2(msg, srcBad2) { - t.Fatal("MAC2 generation/verification failed") - } - } - - checkMAC2([]byte{ - 0x03, 0x31, 0xb9, 0x9e, 0xb0, 0x2a, 0x54, 0xa3, - 0xc1, 0x3f, 0xb4, 0x96, 0x16, 0xb9, 0x25, 0x15, - 0x3d, 0x3a, 0x82, 0xf9, 0x58, 0x36, 0x86, 0x3f, - 0x13, 0x2f, 0xfe, 0xb2, 0x53, 0x20, 0x8c, 0x3f, - 0xba, 0xeb, 0xfb, 0x4b, 0x1b, 0x22, 0x02, 0x69, - 0x2c, 0x90, 0xbc, 0xdc, 0xcf, 0xcf, 0x85, 0xeb, - 0x62, 0x66, 0x6f, 0xe8, 0xe1, 0xa6, 0xa8, 0x4c, - 0xa0, 0x04, 0x23, 0x15, 0x42, 0xac, 0xfa, 0x38, - }) - - checkMAC2([]byte{ - 0x0e, 0x2f, 0x0e, 0xa9, 0x29, 0x03, 0xe1, 0xf3, - 0x24, 0x01, 0x75, 0xad, 0x16, 0xa5, 0x66, 0x85, - 0xca, 0x66, 0xe0, 0xbd, 0xc6, 0x34, 0xd8, 0x84, - 0x09, 0x9a, 0x58, 0x14, 0xfb, 0x05, 0xda, 0xf5, - 0x90, 0xf5, 0x0c, 0x4e, 0x22, 0x10, 0xc9, 0x85, - 0x0f, 0xe3, 0x77, 0x35, 0xe9, 0x6b, 0xc2, 0x55, - 0x32, 0x46, 0xae, 0x25, 0xe0, 0xe3, 0x37, 0x7a, - 0x4b, 0x71, 0xcc, 0xfc, 0x91, 0xdf, 0xd6, 0xca, - 0xfe, 0xee, 0xce, 0x3f, 0x77, 0xa2, 0xfd, 0x59, - 0x8e, 0x73, 0x0a, 0x8d, 0x5c, 0x24, 0x14, 0xca, - 0x38, 0x91, 0xb8, 0x2c, 0x8c, 0xa2, 0x65, 0x7b, - 0xbc, 0x49, 0xbc, 0xb5, 0x58, 0xfc, 0xe3, 0xd7, - 0x02, 0xcf, 0xf7, 0x4c, 0x60, 0x91, 0xed, 0x55, - 0xe9, 0xf9, 0xfe, 0xd1, 0x44, 0x2c, 0x75, 0xf2, - 0xb3, 0x5d, 0x7b, 0x27, 0x56, 0xc0, 0x48, 0x4f, - 0xb0, 0xba, 0xe4, 0x7d, 0xd0, 0xaa, 0xcd, 0x3d, - 0xe3, 0x50, 0xd2, 0xcf, 0xb9, 0xfa, 0x4b, 0x2d, - 0xc6, 0xdf, 0x3b, 0x32, 0x98, 0x45, 0xe6, 0x8f, - 0x1c, 0x5c, 0xa2, 0x20, 0x7d, 0x1c, 0x28, 0xc2, - 0xd4, 0xa1, 0xe0, 0x21, 0x52, 0x8f, 0x1c, 0xd0, - 0x62, 0x97, 0x48, 0xbb, 0xf4, 0xa9, 0xcb, 0x35, - 0xf2, 0x07, 0xd3, 0x50, 0xd8, 0xa9, 0xc5, 0x9a, - 0x0f, 0xbd, 0x37, 0xaf, 0xe1, 0x45, 0x19, 0xee, - 0x41, 0xf3, 0xf7, 0xe5, 0xe0, 0x30, 0x3f, 0xbe, - 0x3d, 0x39, 0x64, 0x00, 0x7a, 0x1a, 0x51, 0x5e, - 0xe1, 0x70, 0x0b, 0xb9, 0x77, 0x5a, 0xf0, 0xc4, - 0x8a, 0xa1, 0x3a, 0x77, 0x1a, 0xe0, 0xc2, 0x06, - 0x91, 0xd5, 0xe9, 0x1c, 0xd3, 0xfe, 0xab, 0x93, - 0x1a, 0x0a, 0x4c, 0xbb, 0xf0, 0xff, 0xdc, 0xaa, - 0x61, 0x73, 0xcb, 0x03, 0x4b, 0x71, 0x68, 0x64, - 0x3d, 0x82, 0x31, 0x41, 0xd7, 0x8b, 0x22, 0x7b, - 0x7d, 0xa1, 0xd5, 0x85, 0x6d, 0xf0, 0x1b, 0xaa, - }) -} diff --git a/device/device_test.go b/device/device_test.go deleted file mode 100644 index 0091e20..0000000 --- a/device/device_test.go +++ /dev/null @@ -1,476 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "bytes" - "encoding/hex" - "fmt" - "io" - "math/rand" - "net/netip" - "os" - "runtime" - "runtime/pprof" - "sync" - "sync/atomic" - "testing" - "time" - - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/conn/bindtest" - "golang.zx2c4.com/wireguard/tun" - "golang.zx2c4.com/wireguard/tun/tuntest" -) - -// uapiCfg returns a string that contains cfg formatted use with IpcSet. -// cfg is a series of alternating key/value strings. -// uapiCfg exists because editors and humans like to insert -// whitespace into configs, which can cause failures, some of which are silent. -// For example, a leading blank newline causes the remainder -// of the config to be silently ignored. -func uapiCfg(cfg ...string) string { - if len(cfg)%2 != 0 { - panic("odd number of args to uapiReader") - } - buf := new(bytes.Buffer) - for i, s := range cfg { - buf.WriteString(s) - sep := byte('\n') - if i%2 == 0 { - sep = '=' - } - buf.WriteByte(sep) - } - return buf.String() -} - -// genConfigs generates a pair of configs that connect to each other. -// The configs use distinct, probably-usable ports. -func genConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) { - var key1, key2 NoisePrivateKey - _, err := rand.Read(key1[:]) - if err != nil { - tb.Errorf("unable to generate private key random bytes: %v", err) - } - _, err = rand.Read(key2[:]) - if err != nil { - tb.Errorf("unable to generate private key random bytes: %v", err) - } - pub1, pub2 := key1.publicKey(), key2.publicKey() - - cfgs[0] = uapiCfg( - "private_key", hex.EncodeToString(key1[:]), - "listen_port", "0", - "replace_peers", "true", - "public_key", hex.EncodeToString(pub2[:]), - "protocol_version", "1", - "replace_allowed_ips", "true", - "allowed_ip", "1.0.0.2/32", - ) - endpointCfgs[0] = uapiCfg( - "public_key", hex.EncodeToString(pub2[:]), - "endpoint", "127.0.0.1:%d", - ) - cfgs[1] = uapiCfg( - "private_key", hex.EncodeToString(key2[:]), - "listen_port", "0", - "replace_peers", "true", - "public_key", hex.EncodeToString(pub1[:]), - "protocol_version", "1", - "replace_allowed_ips", "true", - "allowed_ip", "1.0.0.1/32", - ) - endpointCfgs[1] = uapiCfg( - "public_key", hex.EncodeToString(pub1[:]), - "endpoint", "127.0.0.1:%d", - ) - return -} - -// A testPair is a pair of testPeers. -type testPair [2]testPeer - -// A testPeer is a peer used for testing. -type testPeer struct { - tun *tuntest.ChannelTUN - dev *Device - ip netip.Addr -} - -type SendDirection bool - -const ( - Ping SendDirection = true - Pong SendDirection = false -) - -func (d SendDirection) String() string { - if d == Ping { - return "ping" - } - return "pong" -} - -func (pair *testPair) Send(tb testing.TB, ping SendDirection, done chan struct{}) { - tb.Helper() - p0, p1 := pair[0], pair[1] - if !ping { - // pong is the new ping - p0, p1 = p1, p0 - } - msg := tuntest.Ping(p0.ip, p1.ip) - p1.tun.Outbound <- msg - timer := time.NewTimer(5 * time.Second) - defer timer.Stop() - var err error - select { - case msgRecv := <-p0.tun.Inbound: - if !bytes.Equal(msg, msgRecv) { - err = fmt.Errorf("%s did not transit correctly", ping) - } - case <-timer.C: - err = fmt.Errorf("%s did not transit", ping) - case <-done: - } - if err != nil { - // The error may have occurred because the test is done. - select { - case <-done: - return - default: - } - // Real error. - tb.Error(err) - } -} - -// genTestPair creates a testPair. -func genTestPair(tb testing.TB, realSocket bool) (pair testPair) { - cfg, endpointCfg := genConfigs(tb) - var binds [2]conn.Bind - if realSocket { - binds[0], binds[1] = conn.NewDefaultBind(), conn.NewDefaultBind() - } else { - binds = bindtest.NewChannelBinds() - } - // Bring up a ChannelTun for each config. - for i := range pair { - p := &pair[i] - p.tun = tuntest.NewChannelTUN() - p.ip = netip.AddrFrom4([4]byte{1, 0, 0, byte(i + 1)}) - level := LogLevelVerbose - if _, ok := tb.(*testing.B); ok && !testing.Verbose() { - level = LogLevelError - } - p.dev = NewDevice(p.tun.TUN(), binds[i], NewLogger(level, fmt.Sprintf("dev%d: ", i))) - if err := p.dev.IpcSet(cfg[i]); err != nil { - tb.Errorf("failed to configure device %d: %v", i, err) - p.dev.Close() - continue - } - if err := p.dev.Up(); err != nil { - tb.Errorf("failed to bring up device %d: %v", i, err) - p.dev.Close() - continue - } - endpointCfg[i^1] = fmt.Sprintf(endpointCfg[i^1], p.dev.net.port) - } - for i := range pair { - p := &pair[i] - if err := p.dev.IpcSet(endpointCfg[i]); err != nil { - tb.Errorf("failed to configure device endpoint %d: %v", i, err) - p.dev.Close() - continue - } - // The device is ready. Close it when the test completes. - tb.Cleanup(p.dev.Close) - } - return -} - -func TestTwoDevicePing(t *testing.T) { - goroutineLeakCheck(t) - pair := genTestPair(t, true) - t.Run("ping 1.0.0.1", func(t *testing.T) { - pair.Send(t, Ping, nil) - }) - t.Run("ping 1.0.0.2", func(t *testing.T) { - pair.Send(t, Pong, nil) - }) -} - -func TestUpDown(t *testing.T) { - goroutineLeakCheck(t) - const itrials = 50 - const otrials = 10 - - for n := 0; n < otrials; n++ { - pair := genTestPair(t, false) - for i := range pair { - for k := range pair[i].dev.peers.keyMap { - pair[i].dev.IpcSet(fmt.Sprintf("public_key=%s\npersistent_keepalive_interval=1\n", hex.EncodeToString(k[:]))) - } - } - var wg sync.WaitGroup - wg.Add(len(pair)) - for i := range pair { - go func(d *Device) { - defer wg.Done() - for i := 0; i < itrials; i++ { - if err := d.Up(); err != nil { - t.Errorf("failed up bring up device: %v", err) - } - time.Sleep(time.Duration(rand.Intn(int(time.Nanosecond * (0x10000 - 1))))) - if err := d.Down(); err != nil { - t.Errorf("failed to bring down device: %v", err) - } - time.Sleep(time.Duration(rand.Intn(int(time.Nanosecond * (0x10000 - 1))))) - } - }(pair[i].dev) - } - wg.Wait() - for i := range pair { - pair[i].dev.Up() - pair[i].dev.Close() - } - } -} - -// TestConcurrencySafety does other things concurrently with tunnel use. -// It is intended to be used with the race detector to catch data races. -func TestConcurrencySafety(t *testing.T) { - pair := genTestPair(t, true) - done := make(chan struct{}) - - const warmupIters = 10 - var warmup sync.WaitGroup - warmup.Add(warmupIters) - go func() { - // Send data continuously back and forth until we're done. - // Note that we may continue to attempt to send data - // even after done is closed. - i := warmupIters - for ping := Ping; ; ping = !ping { - pair.Send(t, ping, done) - select { - case <-done: - return - default: - } - if i > 0 { - warmup.Done() - i-- - } - } - }() - warmup.Wait() - - applyCfg := func(cfg string) { - err := pair[0].dev.IpcSet(cfg) - if err != nil { - t.Fatal(err) - } - } - - // Change persistent_keepalive_interval concurrently with tunnel use. - t.Run("persistentKeepaliveInterval", func(t *testing.T) { - var pub NoisePublicKey - for key := range pair[0].dev.peers.keyMap { - pub = key - break - } - cfg := uapiCfg( - "public_key", hex.EncodeToString(pub[:]), - "persistent_keepalive_interval", "1", - ) - for i := 0; i < 1000; i++ { - applyCfg(cfg) - } - }) - - // Change private keys concurrently with tunnel use. - t.Run("privateKey", func(t *testing.T) { - bad := uapiCfg("private_key", "7777777777777777777777777777777777777777777777777777777777777777") - good := uapiCfg("private_key", hex.EncodeToString(pair[0].dev.staticIdentity.privateKey[:])) - // Set iters to a large number like 1000 to flush out data races quickly. - // Don't leave it large. That can cause logical races - // in which the handshake is interleaved with key changes - // such that the private key appears to be unchanging but - // other state gets reset, which can cause handshake failures like - // "Received packet with invalid mac1". - const iters = 1 - for i := 0; i < iters; i++ { - applyCfg(bad) - applyCfg(good) - } - }) - - // Perform bind updates and keepalive sends concurrently with tunnel use. - t.Run("bindUpdate and keepalive", func(t *testing.T) { - const iters = 10 - for i := 0; i < iters; i++ { - for _, peer := range pair { - peer.dev.BindUpdate() - peer.dev.SendKeepalivesToPeersWithCurrentKeypair() - } - } - }) - - close(done) -} - -func BenchmarkLatency(b *testing.B) { - pair := genTestPair(b, true) - - // Establish a connection. - pair.Send(b, Ping, nil) - pair.Send(b, Pong, nil) - - b.ResetTimer() - for i := 0; i < b.N; i++ { - pair.Send(b, Ping, nil) - pair.Send(b, Pong, nil) - } -} - -func BenchmarkThroughput(b *testing.B) { - pair := genTestPair(b, true) - - // Establish a connection. - pair.Send(b, Ping, nil) - pair.Send(b, Pong, nil) - - // Measure how long it takes to receive b.N packets, - // starting when we receive the first packet. - var recv atomic.Uint64 - var elapsed time.Duration - var wg sync.WaitGroup - wg.Add(1) - go func() { - defer wg.Done() - var start time.Time - for { - <-pair[0].tun.Inbound - new := recv.Add(1) - if new == 1 { - start = time.Now() - } - // Careful! Don't change this to else if; b.N can be equal to 1. - if new == uint64(b.N) { - elapsed = time.Since(start) - return - } - } - }() - - // Send packets as fast as we can until we've received enough. - ping := tuntest.Ping(pair[0].ip, pair[1].ip) - pingc := pair[1].tun.Outbound - var sent uint64 - for recv.Load() != uint64(b.N) { - sent++ - pingc <- ping - } - wg.Wait() - - b.ReportMetric(float64(elapsed)/float64(b.N), "ns/op") - b.ReportMetric(1-float64(b.N)/float64(sent), "packet-loss") -} - -func BenchmarkUAPIGet(b *testing.B) { - pair := genTestPair(b, true) - pair.Send(b, Ping, nil) - pair.Send(b, Pong, nil) - b.ReportAllocs() - b.ResetTimer() - for i := 0; i < b.N; i++ { - pair[0].dev.IpcGetOperation(io.Discard) - } -} - -func goroutineLeakCheck(t *testing.T) { - goroutines := func() (int, []byte) { - p := pprof.Lookup("goroutine") - b := new(bytes.Buffer) - p.WriteTo(b, 1) - return p.Count(), b.Bytes() - } - - startGoroutines, startStacks := goroutines() - t.Cleanup(func() { - if t.Failed() { - return - } - // Give goroutines time to exit, if they need it. - for i := 0; i < 10000; i++ { - if runtime.NumGoroutine() <= startGoroutines { - return - } - time.Sleep(1 * time.Millisecond) - } - endGoroutines, endStacks := goroutines() - t.Logf("starting stacks:\n%s\n", startStacks) - t.Logf("ending stacks:\n%s\n", endStacks) - t.Fatalf("expected %d goroutines, got %d, leak?", startGoroutines, endGoroutines) - }) -} - -type fakeBindSized struct { - size int -} - -func (b *fakeBindSized) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err error) { - return nil, 0, nil -} -func (b *fakeBindSized) Close() error { return nil } -func (b *fakeBindSized) SetMark(mark uint32) error { return nil } -func (b *fakeBindSized) Send(bufs [][]byte, ep conn.Endpoint) error { return nil } -func (b *fakeBindSized) ParseEndpoint(s string) (conn.Endpoint, error) { return nil, nil } -func (b *fakeBindSized) BatchSize() int { return b.size } - -type fakeTUNDeviceSized struct { - size int -} - -func (t *fakeTUNDeviceSized) File() *os.File { return nil } -func (t *fakeTUNDeviceSized) Read(bufs [][]byte, sizes []int, offset int) (n int, err error) { - return 0, nil -} -func (t *fakeTUNDeviceSized) Write(bufs [][]byte, offset int) (int, error) { return 0, nil } -func (t *fakeTUNDeviceSized) MTU() (int, error) { return 0, nil } -func (t *fakeTUNDeviceSized) Name() (string, error) { return "", nil } -func (t *fakeTUNDeviceSized) Events() <-chan tun.Event { return nil } -func (t *fakeTUNDeviceSized) Close() error { return nil } -func (t *fakeTUNDeviceSized) BatchSize() int { return t.size } - -func TestBatchSize(t *testing.T) { - d := Device{} - - d.net.bind = &fakeBindSized{1} - d.tun.device = &fakeTUNDeviceSized{1} - if want, got := 1, d.BatchSize(); got != want { - t.Errorf("expected batch size %d, got %d", want, got) - } - - d.net.bind = &fakeBindSized{1} - d.tun.device = &fakeTUNDeviceSized{128} - if want, got := 128, d.BatchSize(); got != want { - t.Errorf("expected batch size %d, got %d", want, got) - } - - d.net.bind = &fakeBindSized{128} - d.tun.device = &fakeTUNDeviceSized{1} - if want, got := 128, d.BatchSize(); got != want { - t.Errorf("expected batch size %d, got %d", want, got) - } - - d.net.bind = &fakeBindSized{128} - d.tun.device = &fakeTUNDeviceSized{128} - if want, got := 128, d.BatchSize(); got != want { - t.Errorf("expected batch size %d, got %d", want, got) - } -} diff --git a/device/endpoint_test.go b/device/endpoint_test.go deleted file mode 100644 index 85482d8..0000000 --- a/device/endpoint_test.go +++ /dev/null @@ -1,49 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "math/rand" - "net/netip" -) - -type DummyEndpoint struct { - src, dst netip.Addr -} - -func CreateDummyEndpoint() (*DummyEndpoint, error) { - var src, dst [16]byte - if _, err := rand.Read(src[:]); err != nil { - return nil, err - } - _, err := rand.Read(dst[:]) - return &DummyEndpoint{netip.AddrFrom16(src), netip.AddrFrom16(dst)}, err -} - -func (e *DummyEndpoint) ClearSrc() {} - -func (e *DummyEndpoint) SrcToString() string { - return netip.AddrPortFrom(e.SrcIP(), 1000).String() -} - -func (e *DummyEndpoint) DstToString() string { - return netip.AddrPortFrom(e.DstIP(), 1000).String() -} - -func (e *DummyEndpoint) DstToBytes() []byte { - out := e.DstIP().AsSlice() - out = append(out, byte(1000&0xff)) - out = append(out, byte((1000>>8)&0xff)) - return out -} - -func (e *DummyEndpoint) DstIP() netip.Addr { - return e.dst -} - -func (e *DummyEndpoint) SrcIP() netip.Addr { - return e.src -} diff --git a/device/kdf_test.go b/device/kdf_test.go deleted file mode 100644 index 325db59..0000000 --- a/device/kdf_test.go +++ /dev/null @@ -1,85 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "encoding/hex" - "testing" - - "golang.org/x/crypto/blake2s" -) - -type KDFTest struct { - key string - input string - t0 string - t1 string - t2 string -} - -func assertEquals(t *testing.T, a, b string) { - if a != b { - t.Fatal("expected", a, "=", b) - } -} - -func TestKDF(t *testing.T) { - tests := []KDFTest{ - { - key: "746573742d6b6579", - input: "746573742d696e707574", - t0: "6f0e5ad38daba1bea8a0d213688736f19763239305e0f58aba697f9ffc41c633", - t1: "df1194df20802a4fe594cde27e92991c8cae66c366e8106aaa937a55fa371e8a", - t2: "fac6e2745a325f5dc5d11a5b165aad08b0ada28e7b4e666b7c077934a4d76c24", - }, - { - key: "776972656775617264", - input: "776972656775617264", - t0: "491d43bbfdaa8750aaf535e334ecbfe5129967cd64635101c566d4caefda96e8", - t1: "1e71a379baefd8a79aa4662212fcafe19a23e2b609a3db7d6bcba8f560e3d25f", - t2: "31e1ae48bddfbe5de38f295e5452b1909a1b4e38e183926af3780b0c1e1f0160", - }, - { - key: "", - input: "", - t0: "8387b46bf43eccfcf349552a095d8315c4055beb90208fb1be23b894bc2ed5d0", - t1: "58a0e5f6faefccf4807bff1f05fa8a9217945762040bcec2f4b4a62bdfe0e86e", - t2: "0ce6ea98ec548f8e281e93e32db65621c45eb18dc6f0a7ad94178610a2f7338e", - }, - } - - var t0, t1, t2 [blake2s.Size]byte - - for _, test := range tests { - key, _ := hex.DecodeString(test.key) - input, _ := hex.DecodeString(test.input) - KDF3(&t0, &t1, &t2, key, input) - t0s := hex.EncodeToString(t0[:]) - t1s := hex.EncodeToString(t1[:]) - t2s := hex.EncodeToString(t2[:]) - assertEquals(t, t0s, test.t0) - assertEquals(t, t1s, test.t1) - assertEquals(t, t2s, test.t2) - } - - for _, test := range tests { - key, _ := hex.DecodeString(test.key) - input, _ := hex.DecodeString(test.input) - KDF2(&t0, &t1, key, input) - t0s := hex.EncodeToString(t0[:]) - t1s := hex.EncodeToString(t1[:]) - assertEquals(t, t0s, test.t0) - assertEquals(t, t1s, test.t1) - } - - for _, test := range tests { - key, _ := hex.DecodeString(test.key) - input, _ := hex.DecodeString(test.input) - KDF1(&t0, key, input) - t0s := hex.EncodeToString(t0[:]) - assertEquals(t, t0s, test.t0) - } -} diff --git a/device/noise_test.go b/device/noise_test.go deleted file mode 100644 index f0928ac..0000000 --- a/device/noise_test.go +++ /dev/null @@ -1,179 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "bytes" - "encoding/binary" - "testing" - - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/tun/tuntest" -) - -func TestCurveWrappers(t *testing.T) { - sk1, err := newPrivateKey() - assertNil(t, err) - - sk2, err := newPrivateKey() - assertNil(t, err) - - pk1 := sk1.publicKey() - pk2 := sk2.publicKey() - - ss1, err1 := sk1.sharedSecret(pk2) - ss2, err2 := sk2.sharedSecret(pk1) - - if ss1 != ss2 || err1 != nil || err2 != nil { - t.Fatal("Failed to compute shared secet") - } -} - -func randDevice(t *testing.T) *Device { - sk, err := newPrivateKey() - if err != nil { - t.Fatal(err) - } - tun := tuntest.NewChannelTUN() - logger := NewLogger(LogLevelError, "") - device := NewDevice(tun.TUN(), conn.NewDefaultBind(), logger) - device.SetPrivateKey(sk) - return device -} - -func assertNil(t *testing.T, err error) { - if err != nil { - t.Fatal(err) - } -} - -func assertEqual(t *testing.T, a, b []byte) { - if !bytes.Equal(a, b) { - t.Fatal(a, "!=", b) - } -} - -func TestNoiseHandshake(t *testing.T) { - dev1 := randDevice(t) - dev2 := randDevice(t) - - defer dev1.Close() - defer dev2.Close() - - 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) - } - peer1.Start() - peer2.Start() - - assertEqual( - t, - peer1.handshake.precomputedStaticStatic[:], - peer2.handshake.precomputedStaticStatic[:], - ) - - /* simulate handshake */ - - // initiation message - - t.Log("exchange initiation message") - - msg1, err := dev1.CreateMessageInitiation(peer2) - assertNil(t, err) - - packet := make([]byte, 0, 256) - writer := bytes.NewBuffer(packet) - err = binary.Write(writer, binary.LittleEndian, msg1) - assertNil(t, err) - peer := dev2.ConsumeMessageInitiation(msg1) - if peer == nil { - t.Fatal("handshake failed at initiation message") - } - - assertEqual( - t, - peer1.handshake.chainKey[:], - peer2.handshake.chainKey[:], - ) - - assertEqual( - t, - peer1.handshake.hash[:], - peer2.handshake.hash[:], - ) - - // response message - - t.Log("exchange response message") - - msg2, err := dev2.CreateMessageResponse(peer1) - assertNil(t, err) - - peer = dev1.ConsumeMessageResponse(msg2) - if peer == nil { - t.Fatal("handshake failed at response message") - } - - assertEqual( - t, - peer1.handshake.chainKey[:], - peer2.handshake.chainKey[:], - ) - - assertEqual( - t, - peer1.handshake.hash[:], - peer2.handshake.hash[:], - ) - - // key pairs - - t.Log("deriving keys") - - err = peer1.BeginSymmetricSession() - if err != nil { - t.Fatal("failed to derive keypair for peer 1", err) - } - - err = peer2.BeginSymmetricSession() - if err != nil { - t.Fatal("failed to derive keypair for peer 2", err) - } - - key1 := peer1.keypairs.next.Load() - key2 := peer2.keypairs.current - - // encrypting / decryption test - - t.Log("test key pairs") - - func() { - testMsg := []byte("wireguard test message 1") - var err error - var out []byte - var nonce [12]byte - out = key1.send.Seal(out, nonce[:], testMsg, nil) - out, err = key2.receive.Open(out[:0], nonce[:], out, nil) - assertNil(t, err) - assertEqual(t, out, testMsg) - }() - - func() { - testMsg := []byte("wireguard test message 2") - var err error - var out []byte - var nonce [12]byte - out = key2.send.Seal(out, nonce[:], testMsg, nil) - out, err = key1.receive.Open(out[:0], nonce[:], out, nil) - assertNil(t, err) - assertEqual(t, out, testMsg) - }() -} diff --git a/device/pools_test.go b/device/pools_test.go deleted file mode 100644 index 8381d5a..0000000 --- a/device/pools_test.go +++ /dev/null @@ -1,141 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -import ( - "math/rand" - "runtime" - "sync" - "sync/atomic" - "testing" - "time" -) - -func TestWaitPool(t *testing.T) { - t.Skip("Currently disabled") - var wg sync.WaitGroup - var trials atomic.Int32 - startTrials := int32(100000) - if raceEnabled { - // This test can be very slow with -race. - startTrials /= 10 - } - trials.Store(startTrials) - workers := runtime.NumCPU() + 2 - if workers-4 <= 0 { - t.Skip("Not enough cores") - } - p := NewWaitPool(uint32(workers-4), func() any { return make([]byte, 16) }) - wg.Add(workers) - var max atomic.Uint32 - updateMax := func() { - p.lock.Lock() - count := p.count - p.lock.Unlock() - if count > p.max { - t.Errorf("count (%d) > max (%d)", count, p.max) - } - for { - old := max.Load() - if count <= old { - break - } - if max.CompareAndSwap(old, count) { - break - } - } - } - for i := 0; i < workers; i++ { - go func() { - defer wg.Done() - for trials.Add(-1) > 0 { - updateMax() - x := p.Get() - updateMax() - time.Sleep(time.Duration(rand.Intn(100)) * time.Microsecond) - updateMax() - p.Put(x) - updateMax() - } - }() - } - wg.Wait() - if max.Load() != p.max { - t.Errorf("Actual maximum count (%d) != ideal maximum count (%d)", max, p.max) - } -} - -func BenchmarkWaitPool(b *testing.B) { - var wg sync.WaitGroup - var trials atomic.Int32 - trials.Store(int32(b.N)) - workers := runtime.NumCPU() + 2 - if workers-4 <= 0 { - b.Skip("Not enough cores") - } - p := NewWaitPool(uint32(workers-4), func() any { return make([]byte, 16) }) - wg.Add(workers) - b.ResetTimer() - for i := 0; i < workers; i++ { - go func() { - defer wg.Done() - for trials.Add(-1) > 0 { - x := p.Get() - time.Sleep(time.Duration(rand.Intn(100)) * time.Microsecond) - p.Put(x) - } - }() - } - wg.Wait() -} - -func BenchmarkWaitPoolEmpty(b *testing.B) { - var wg sync.WaitGroup - var trials atomic.Int32 - trials.Store(int32(b.N)) - workers := runtime.NumCPU() + 2 - if workers-4 <= 0 { - b.Skip("Not enough cores") - } - p := NewWaitPool(0, func() any { return make([]byte, 16) }) - wg.Add(workers) - b.ResetTimer() - for i := 0; i < workers; i++ { - go func() { - defer wg.Done() - for trials.Add(-1) > 0 { - x := p.Get() - time.Sleep(time.Duration(rand.Intn(100)) * time.Microsecond) - p.Put(x) - } - }() - } - wg.Wait() -} - -func BenchmarkSyncPool(b *testing.B) { - var wg sync.WaitGroup - var trials atomic.Int32 - trials.Store(int32(b.N)) - workers := runtime.NumCPU() + 2 - if workers-4 <= 0 { - b.Skip("Not enough cores") - } - p := sync.Pool{New: func() any { return make([]byte, 16) }} - wg.Add(workers) - b.ResetTimer() - for i := 0; i < workers; i++ { - go func() { - defer wg.Done() - for trials.Add(-1) > 0 { - x := p.Get() - time.Sleep(time.Duration(rand.Intn(100)) * time.Microsecond) - p.Put(x) - } - }() - } - wg.Wait() -} diff --git a/device/race_disabled_test.go b/device/race_disabled_test.go deleted file mode 100644 index 14b3284..0000000 --- a/device/race_disabled_test.go +++ /dev/null @@ -1,10 +0,0 @@ -//go:build !race - -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -const raceEnabled = false diff --git a/device/race_enabled_test.go b/device/race_enabled_test.go deleted file mode 100644 index f1ea5cf..0000000 --- a/device/race_enabled_test.go +++ /dev/null @@ -1,10 +0,0 @@ -//go:build race - -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package device - -const raceEnabled = true diff --git a/format_test.go b/format_test.go deleted file mode 100644 index 4d02c48..0000000 --- a/format_test.go +++ /dev/null @@ -1,51 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ -package main - -import ( - "bytes" - "go/format" - "io/fs" - "os" - "path/filepath" - "runtime" - "sync" - "testing" -) - -func TestFormatting(t *testing.T) { - var wg sync.WaitGroup - filepath.WalkDir(".", func(path string, d fs.DirEntry, err error) error { - if err != nil { - t.Errorf("unable to walk %s: %v", path, err) - return nil - } - if d.IsDir() || filepath.Ext(path) != ".go" { - return nil - } - wg.Add(1) - go func(path string) { - defer wg.Done() - src, err := os.ReadFile(path) - if err != nil { - t.Errorf("unable to read %s: %v", path, err) - return - } - if runtime.GOOS == "windows" { - src = bytes.ReplaceAll(src, []byte{'\r', '\n'}, []byte{'\n'}) - } - formatted, err := format.Source(src) - if err != nil { - t.Errorf("unable to format %s: %v", path, err) - return - } - if !bytes.Equal(src, formatted) { - t.Errorf("unformatted code: %s", path) - } - }(path) - return nil - }) - wg.Wait() -} diff --git a/go.mod b/go.mod index 2a80e00..85766cf 100644 --- a/go.mod +++ b/go.mod @@ -7,10 +7,4 @@ require ( golang.org/x/net v0.39.0 golang.org/x/sys v0.32.0 golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 - gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c -) - -require ( - 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 61875c1..c9da2b3 100644 --- a/go.sum +++ b/go.sum @@ -1,14 +1,8 @@ -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= golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc= golang.org/x/net v0.39.0 h1:ZCu7HMWDxpXpaiKdhzIfaltL9Lp31x/3fCP11bc6/fY= golang.org/x/net v0.39.0/go.mod h1:X7NRbYVEA+ewNkCNyJ513WmMdQ3BineSwVtN2zD/d+E= golang.org/x/sys v0.32.0 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20= golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= -golang.org/x/time v0.7.0 h1:ntUhktv3OPE6TgYxXWv9vKvUSJyIFJlyohwbkEwPrKQ= -golang.org/x/time v0.7.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI= -gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c h1:m/r7OM+Y2Ty1sgBQ7Qb27VgIMBW8ZZhT4gLnUyDIhzI= -gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g= diff --git a/ipc/namedpipe/namedpipe_test.go b/ipc/namedpipe/namedpipe_test.go deleted file mode 100644 index 998453b..0000000 --- a/ipc/namedpipe/namedpipe_test.go +++ /dev/null @@ -1,674 +0,0 @@ -// Copyright 2021 The Go Authors. All rights reserved. -// Copyright 2015 Microsoft -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -//go:build windows - -package namedpipe_test - -import ( - "bufio" - "bytes" - "context" - "errors" - "io" - "net" - "os" - "sync" - "syscall" - "testing" - "time" - - "golang.org/x/sys/windows" - "golang.zx2c4.com/wireguard/ipc/namedpipe" -) - -func randomPipePath() string { - guid, err := windows.GenerateGUID() - if err != nil { - panic(err) - } - return `\\.\PIPE\go-namedpipe-test-` + guid.String() -} - -func TestPingPong(t *testing.T) { - const ( - ping = 42 - pong = 24 - ) - pipePath := randomPipePath() - listener, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatalf("unable to listen on pipe: %v", err) - } - defer listener.Close() - go func() { - incoming, err := listener.Accept() - if err != nil { - t.Fatalf("unable to accept pipe connection: %v", err) - } - defer incoming.Close() - var data [1]byte - _, err = incoming.Read(data[:]) - if err != nil { - t.Fatalf("unable to read ping from pipe: %v", err) - } - if data[0] != ping { - t.Fatalf("expected ping, got %d", data[0]) - } - data[0] = pong - _, err = incoming.Write(data[:]) - if err != nil { - t.Fatalf("unable to write pong to pipe: %v", err) - } - }() - client, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err != nil { - t.Fatalf("unable to dial pipe: %v", err) - } - defer client.Close() - client.SetDeadline(time.Now().Add(time.Second * 5)) - var data [1]byte - data[0] = ping - _, err = client.Write(data[:]) - if err != nil { - t.Fatalf("unable to write ping to pipe: %v", err) - } - _, err = client.Read(data[:]) - if err != nil { - t.Fatalf("unable to read pong from pipe: %v", err) - } - if data[0] != pong { - t.Fatalf("expected pong, got %d", data[0]) - } -} - -func TestDialUnknownFailsImmediately(t *testing.T) { - _, err := namedpipe.DialTimeout(randomPipePath(), time.Duration(0)) - if !errors.Is(err, syscall.ENOENT) { - t.Fatalf("expected ENOENT got %v", err) - } -} - -func TestDialListenerTimesOut(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - pipe, err := namedpipe.DialTimeout(pipePath, 10*time.Millisecond) - if err == nil { - pipe.Close() - } - if err != os.ErrDeadlineExceeded { - t.Fatalf("expected os.ErrDeadlineExceeded, got %v", err) - } -} - -func TestDialContextListenerTimesOut(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - d := 10 * time.Millisecond - ctx, _ := context.WithTimeout(context.Background(), d) - pipe, err := namedpipe.DialContext(ctx, pipePath) - if err == nil { - pipe.Close() - } - if err != context.DeadlineExceeded { - t.Fatalf("expected context.DeadlineExceeded, got %v", err) - } -} - -func TestDialListenerGetsCancelled(t *testing.T) { - pipePath := randomPipePath() - ctx, cancel := context.WithCancel(context.Background()) - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - ch := make(chan error) - go func(ctx context.Context, ch chan error) { - _, err := namedpipe.DialContext(ctx, pipePath) - ch <- err - }(ctx, ch) - time.Sleep(time.Millisecond * 30) - cancel() - err = <-ch - if err != context.Canceled { - t.Fatalf("expected context.Canceled, got %v", err) - } -} - -func TestDialAccessDeniedWithRestrictedSD(t *testing.T) { - if windows.NewLazySystemDLL("ntdll.dll").NewProc("wine_get_version").Find() == nil { - t.Skip("dacls on named pipes are broken on wine") - } - pipePath := randomPipePath() - sd, _ := windows.SecurityDescriptorFromString("D:") - l, err := (&namedpipe.ListenConfig{ - SecurityDescriptor: sd, - }).Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - pipe, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err == nil { - pipe.Close() - } - if !errors.Is(err, windows.ERROR_ACCESS_DENIED) { - t.Fatalf("expected ERROR_ACCESS_DENIED, got %v", err) - } -} - -func getConnection(cfg *namedpipe.ListenConfig) (client, server net.Conn, err error) { - pipePath := randomPipePath() - if cfg == nil { - cfg = &namedpipe.ListenConfig{} - } - l, err := cfg.Listen(pipePath) - if err != nil { - return - } - defer l.Close() - - type response struct { - c net.Conn - err error - } - ch := make(chan response) - go func() { - c, err := l.Accept() - ch <- response{c, err} - }() - - c, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err != nil { - return - } - - r := <-ch - if err = r.err; err != nil { - c.Close() - return - } - - client = c - server = r.c - return -} - -func TestReadTimeout(t *testing.T) { - c, s, err := getConnection(nil) - if err != nil { - t.Fatal(err) - } - defer c.Close() - defer s.Close() - - c.SetReadDeadline(time.Now().Add(10 * time.Millisecond)) - - buf := make([]byte, 10) - _, err = c.Read(buf) - if err != os.ErrDeadlineExceeded { - t.Fatalf("expected os.ErrDeadlineExceeded, got %v", err) - } -} - -func server(l net.Listener, ch chan int) { - c, err := l.Accept() - if err != nil { - panic(err) - } - rw := bufio.NewReadWriter(bufio.NewReader(c), bufio.NewWriter(c)) - s, err := rw.ReadString('\n') - if err != nil { - panic(err) - } - _, err = rw.WriteString("got " + s) - if err != nil { - panic(err) - } - err = rw.Flush() - if err != nil { - panic(err) - } - c.Close() - ch <- 1 -} - -func TestFullListenDialReadWrite(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - - ch := make(chan int) - go server(l, ch) - - c, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err != nil { - t.Fatal(err) - } - defer c.Close() - - rw := bufio.NewReadWriter(bufio.NewReader(c), bufio.NewWriter(c)) - _, err = rw.WriteString("hello world\n") - if err != nil { - t.Fatal(err) - } - err = rw.Flush() - if err != nil { - t.Fatal(err) - } - - s, err := rw.ReadString('\n') - if err != nil { - t.Fatal(err) - } - ms := "got hello world\n" - if s != ms { - t.Errorf("expected '%s', got '%s'", ms, s) - } - - <-ch -} - -func TestCloseAbortsListen(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - - ch := make(chan error) - go func() { - _, err := l.Accept() - ch <- err - }() - - time.Sleep(30 * time.Millisecond) - l.Close() - - err = <-ch - if err != net.ErrClosed { - t.Fatalf("expected net.ErrClosed, got %v", err) - } -} - -func ensureEOFOnClose(t *testing.T, r io.Reader, w io.Closer) { - b := make([]byte, 10) - w.Close() - n, err := r.Read(b) - if n > 0 { - t.Errorf("unexpected byte count %d", n) - } - if err != io.EOF { - t.Errorf("expected EOF: %v", err) - } -} - -func TestCloseClientEOFServer(t *testing.T) { - c, s, err := getConnection(nil) - if err != nil { - t.Fatal(err) - } - defer c.Close() - defer s.Close() - ensureEOFOnClose(t, c, s) -} - -func TestCloseServerEOFClient(t *testing.T) { - c, s, err := getConnection(nil) - if err != nil { - t.Fatal(err) - } - defer c.Close() - defer s.Close() - ensureEOFOnClose(t, s, c) -} - -func TestCloseWriteEOF(t *testing.T) { - cfg := &namedpipe.ListenConfig{ - MessageMode: true, - } - c, s, err := getConnection(cfg) - if err != nil { - t.Fatal(err) - } - defer c.Close() - defer s.Close() - - type closeWriter interface { - CloseWrite() error - } - - err = c.(closeWriter).CloseWrite() - if err != nil { - t.Fatal(err) - } - - b := make([]byte, 10) - _, err = s.Read(b) - if err != io.EOF { - t.Fatal(err) - } -} - -func TestAcceptAfterCloseFails(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - l.Close() - _, err = l.Accept() - if err != net.ErrClosed { - t.Fatalf("expected net.ErrClosed, got %v", err) - } -} - -func TestDialTimesOutByDefault(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - pipe, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) // Should timeout after 2 seconds. - if err == nil { - pipe.Close() - } - if err != os.ErrDeadlineExceeded { - t.Fatalf("expected os.ErrDeadlineExceeded, got %v", err) - } -} - -func TestTimeoutPendingRead(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - - serverDone := make(chan struct{}) - - go func() { - s, err := l.Accept() - if err != nil { - t.Fatal(err) - } - time.Sleep(1 * time.Second) - s.Close() - close(serverDone) - }() - - client, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err != nil { - t.Fatal(err) - } - defer client.Close() - - clientErr := make(chan error) - go func() { - buf := make([]byte, 10) - _, err = client.Read(buf) - clientErr <- err - }() - - time.Sleep(100 * time.Millisecond) // make *sure* the pipe is reading before we set the deadline - client.SetReadDeadline(time.Unix(1, 0)) - - select { - case err = <-clientErr: - if err != os.ErrDeadlineExceeded { - t.Fatalf("expected os.ErrDeadlineExceeded, got %v", err) - } - case <-time.After(100 * time.Millisecond): - t.Fatalf("timed out while waiting for read to cancel") - <-clientErr - } - <-serverDone -} - -func TestTimeoutPendingWrite(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - - serverDone := make(chan struct{}) - - go func() { - s, err := l.Accept() - if err != nil { - t.Fatal(err) - } - time.Sleep(1 * time.Second) - s.Close() - close(serverDone) - }() - - client, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err != nil { - t.Fatal(err) - } - defer client.Close() - - clientErr := make(chan error) - go func() { - _, err = client.Write([]byte("this should timeout")) - clientErr <- err - }() - - time.Sleep(100 * time.Millisecond) // make *sure* the pipe is writing before we set the deadline - client.SetWriteDeadline(time.Unix(1, 0)) - - select { - case err = <-clientErr: - if err != os.ErrDeadlineExceeded { - t.Fatalf("expected os.ErrDeadlineExceeded, got %v", err) - } - case <-time.After(100 * time.Millisecond): - t.Fatalf("timed out while waiting for write to cancel") - <-clientErr - } - <-serverDone -} - -type CloseWriter interface { - CloseWrite() error -} - -func TestEchoWithMessaging(t *testing.T) { - pipePath := randomPipePath() - l, err := (&namedpipe.ListenConfig{ - MessageMode: true, // Use message mode so that CloseWrite() is supported - InputBufferSize: 65536, // Use 64KB buffers to improve performance - OutputBufferSize: 65536, - }).Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - - listenerDone := make(chan bool) - clientDone := make(chan bool) - go func() { - // server echo - conn, err := l.Accept() - if err != nil { - t.Fatal(err) - } - defer conn.Close() - - time.Sleep(500 * time.Millisecond) // make *sure* we don't begin to read before eof signal is sent - _, err = io.Copy(conn, conn) - if err != nil { - t.Fatal(err) - } - conn.(CloseWriter).CloseWrite() - close(listenerDone) - }() - client, err := namedpipe.DialTimeout(pipePath, time.Second) - if err != nil { - t.Fatal(err) - } - defer client.Close() - - go func() { - // client read back - bytes := make([]byte, 2) - n, e := client.Read(bytes) - if e != nil { - t.Fatal(e) - } - if n != 2 || bytes[0] != 0 || bytes[1] != 1 { - t.Fatalf("expected 2 bytes, got %v", n) - } - close(clientDone) - }() - - payload := make([]byte, 2) - payload[0] = 0 - payload[1] = 1 - - n, err := client.Write(payload) - if err != nil { - t.Fatal(err) - } - if n != 2 { - t.Fatalf("expected 2 bytes, got %v", n) - } - client.(CloseWriter).CloseWrite() - <-listenerDone - <-clientDone -} - -func TestConnectRace(t *testing.T) { - pipePath := randomPipePath() - l, err := namedpipe.Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - go func() { - for { - s, err := l.Accept() - if err == net.ErrClosed { - return - } - - if err != nil { - t.Fatal(err) - } - s.Close() - } - }() - - for i := 0; i < 1000; i++ { - c, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err != nil { - t.Fatal(err) - } - c.Close() - } -} - -func TestMessageReadMode(t *testing.T) { - if maj, _, _ := windows.RtlGetNtVersionNumbers(); maj <= 8 { - t.Skipf("Skipping on Windows %d", maj) - } - var wg sync.WaitGroup - defer wg.Wait() - pipePath := randomPipePath() - l, err := (&namedpipe.ListenConfig{MessageMode: true}).Listen(pipePath) - if err != nil { - t.Fatal(err) - } - defer l.Close() - - msg := ([]byte)("hello world") - - wg.Add(1) - go func() { - defer wg.Done() - s, err := l.Accept() - if err != nil { - t.Fatal(err) - } - _, err = s.Write(msg) - if err != nil { - t.Fatal(err) - } - s.Close() - }() - - c, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err != nil { - t.Fatal(err) - } - defer c.Close() - - mode := uint32(windows.PIPE_READMODE_MESSAGE) - err = windows.SetNamedPipeHandleState(c.(interface{ Handle() windows.Handle }).Handle(), &mode, nil, nil) - if err != nil { - t.Fatal(err) - } - - ch := make([]byte, 1) - var vmsg []byte - for { - n, err := c.Read(ch) - if err == io.EOF { - break - } - if err != nil { - t.Fatal(err) - } - if n != 1 { - t.Fatalf("expected 1, got %d", n) - } - vmsg = append(vmsg, ch[0]) - } - if !bytes.Equal(msg, vmsg) { - t.Fatalf("expected %s, got %s", msg, vmsg) - } -} - -func TestListenConnectRace(t *testing.T) { - if testing.Short() { - t.Skip("Skipping long race test") - } - pipePath := randomPipePath() - for i := 0; i < 50 && !t.Failed(); i++ { - var wg sync.WaitGroup - wg.Add(1) - go func() { - c, err := namedpipe.DialTimeout(pipePath, time.Duration(0)) - if err == nil { - c.Close() - } - wg.Done() - }() - s, err := namedpipe.Listen(pipePath) - if err != nil { - t.Error(i, err) - } else { - s.Close() - } - wg.Wait() - } -} diff --git a/main.go b/main.go deleted file mode 100644 index b6989e2..0000000 --- a/main.go +++ /dev/null @@ -1,268 +0,0 @@ -//go:build !windows - -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package main - -import ( - "fmt" - "os" - "os/signal" - "runtime" - "strconv" - - "golang.org/x/sys/unix" - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/device" - "golang.zx2c4.com/wireguard/ipc" - "golang.zx2c4.com/wireguard/tun" -) - -const ( - ExitSetupSuccess = 0 - ExitSetupFailed = 1 -) - -const ( - ENV_WG_TUN_FD = "WG_TUN_FD" - ENV_WG_UAPI_FD = "WG_UAPI_FD" - ENV_WG_PROCESS_FOREGROUND = "WG_PROCESS_FOREGROUND" -) - -func printUsage() { - fmt.Printf("Usage: %s [-f/--foreground] INTERFACE-NAME\n", os.Args[0]) -} - -func warning() { - switch runtime.GOOS { - case "linux", "freebsd", "openbsd": - if os.Getenv(ENV_WG_PROCESS_FOREGROUND) == "1" { - return - } - default: - return - } - - fmt.Fprintln(os.Stderr, "┌──────────────────────────────────────────────────────┐") - fmt.Fprintln(os.Stderr, "│ │") - fmt.Fprintln(os.Stderr, "│ Running wireguard-go is not required because this │") - fmt.Fprintln(os.Stderr, "│ kernel has first class support for WireGuard. For │") - fmt.Fprintln(os.Stderr, "│ information on installing the kernel module, │") - fmt.Fprintln(os.Stderr, "│ please visit: │") - fmt.Fprintln(os.Stderr, "│ https://www.wireguard.com/install/ │") - fmt.Fprintln(os.Stderr, "│ │") - fmt.Fprintln(os.Stderr, "└──────────────────────────────────────────────────────┘") -} - -func main() { - if len(os.Args) == 2 && os.Args[1] == "--version" { - fmt.Printf("wireguard-go v%s\n\nUserspace WireGuard daemon for %s-%s.\nInformation available at https://www.wireguard.com.\nCopyright (C) Jason A. Donenfeld .\n", Version, runtime.GOOS, runtime.GOARCH) - return - } - - warning() - - var foreground bool - var interfaceName string - if len(os.Args) < 2 || len(os.Args) > 3 { - printUsage() - return - } - - switch os.Args[1] { - - case "-f", "--foreground": - foreground = true - if len(os.Args) != 3 { - printUsage() - return - } - interfaceName = os.Args[2] - - default: - foreground = false - if len(os.Args) != 2 { - printUsage() - return - } - interfaceName = os.Args[1] - } - - if !foreground { - foreground = os.Getenv(ENV_WG_PROCESS_FOREGROUND) == "1" - } - - // get log level (default: info) - - logLevel := func() int { - switch os.Getenv("LOG_LEVEL") { - case "verbose", "debug": - return device.LogLevelVerbose - case "error": - return device.LogLevelError - case "silent": - return device.LogLevelSilent - } - return device.LogLevelError - }() - - // open TUN device (or use supplied fd) - - tdev, err := func() (tun.Device, error) { - tunFdStr := os.Getenv(ENV_WG_TUN_FD) - if tunFdStr == "" { - return tun.CreateTUN(interfaceName, device.DefaultMTU) - } - - // construct tun device from supplied fd - - fd, err := strconv.ParseUint(tunFdStr, 10, 32) - if err != nil { - return nil, err - } - - err = unix.SetNonblock(int(fd), true) - if err != nil { - return nil, err - } - - file := os.NewFile(uintptr(fd), "") - return tun.CreateTUNFromFile(file, device.DefaultMTU) - }() - - if err == nil { - realInterfaceName, err2 := tdev.Name() - if err2 == nil { - interfaceName = realInterfaceName - } - } - - logger := device.NewLogger( - logLevel, - fmt.Sprintf("(%s) ", interfaceName), - ) - - logger.Verbosef("Starting wireguard-go version %s", Version) - - if err != nil { - logger.Errorf("Failed to create TUN device: %v", err) - os.Exit(ExitSetupFailed) - } - - // open UAPI file (or use supplied fd) - - fileUAPI, err := func() (*os.File, error) { - uapiFdStr := os.Getenv(ENV_WG_UAPI_FD) - if uapiFdStr == "" { - return ipc.UAPIOpen(interfaceName) - } - - // use supplied fd - - fd, err := strconv.ParseUint(uapiFdStr, 10, 32) - if err != nil { - return nil, err - } - - return os.NewFile(uintptr(fd), ""), nil - }() - if err != nil { - logger.Errorf("UAPI listen error: %v", err) - os.Exit(ExitSetupFailed) - return - } - // daemonize the process - - if !foreground { - env := os.Environ() - env = append(env, fmt.Sprintf("%s=3", ENV_WG_TUN_FD)) - env = append(env, fmt.Sprintf("%s=4", ENV_WG_UAPI_FD)) - env = append(env, fmt.Sprintf("%s=1", ENV_WG_PROCESS_FOREGROUND)) - files := [3]*os.File{} - if os.Getenv("LOG_LEVEL") != "" && logLevel != device.LogLevelSilent { - files[0], _ = os.Open(os.DevNull) - files[1] = os.Stdout - files[2] = os.Stderr - } else { - files[0], _ = os.Open(os.DevNull) - files[1], _ = os.Open(os.DevNull) - files[2], _ = os.Open(os.DevNull) - } - attr := &os.ProcAttr{ - Files: []*os.File{ - files[0], // stdin - files[1], // stdout - files[2], // stderr - tdev.File(), - fileUAPI, - }, - Dir: ".", - Env: env, - } - - path, err := os.Executable() - if err != nil { - logger.Errorf("Failed to determine executable: %v", err) - os.Exit(ExitSetupFailed) - } - - process, err := os.StartProcess( - path, - os.Args, - attr, - ) - if err != nil { - logger.Errorf("Failed to daemonize: %v", err) - os.Exit(ExitSetupFailed) - } - process.Release() - return - } - - device := device.NewDevice(tdev, conn.NewDefaultBind(), logger) - - logger.Verbosef("Device started") - - errs := make(chan error) - term := make(chan os.Signal, 1) - - uapi, err := ipc.UAPIListen(interfaceName, fileUAPI) - if err != nil { - logger.Errorf("Failed to listen on uapi socket: %v", err) - os.Exit(ExitSetupFailed) - } - - go func() { - for { - conn, err := uapi.Accept() - if err != nil { - errs <- err - return - } - go device.IpcHandle(conn) - } - }() - - logger.Verbosef("UAPI listener started") - - // wait for program to terminate - - signal.Notify(term, unix.SIGTERM) - signal.Notify(term, os.Interrupt) - - select { - case <-term: - case <-errs: - case <-device.Wait(): - } - - // clean up - - uapi.Close() - device.Close() - - logger.Verbosef("Shutting down") -} diff --git a/main_windows.go b/main_windows.go deleted file mode 100644 index 67036cf..0000000 --- a/main_windows.go +++ /dev/null @@ -1,99 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package main - -import ( - "fmt" - "os" - "os/signal" - - "golang.org/x/sys/windows" - - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/device" - "golang.zx2c4.com/wireguard/ipc" - - "golang.zx2c4.com/wireguard/tun" -) - -const ( - ExitSetupSuccess = 0 - ExitSetupFailed = 1 -) - -func main() { - if len(os.Args) != 2 { - os.Exit(ExitSetupFailed) - } - interfaceName := os.Args[1] - - fmt.Fprintln(os.Stderr, "Warning: this is a test program for Windows, mainly used for debugging this Go package. For a real WireGuard for Windows client, the repo you want is , which includes this code as a module.") - - logger := device.NewLogger( - device.LogLevelVerbose, - fmt.Sprintf("(%s) ", interfaceName), - ) - logger.Verbosef("Starting wireguard-go version %s", Version) - - tun, err := tun.CreateTUN(interfaceName, 0) - if err == nil { - realInterfaceName, err2 := tun.Name() - if err2 == nil { - interfaceName = realInterfaceName - } - } else { - logger.Errorf("Failed to create TUN device: %v", err) - os.Exit(ExitSetupFailed) - } - - device := device.NewDevice(tun, conn.NewDefaultBind(), logger) - err = device.Up() - if err != nil { - logger.Errorf("Failed to bring up device: %v", err) - os.Exit(ExitSetupFailed) - } - logger.Verbosef("Device started") - - uapi, err := ipc.UAPIListen(interfaceName) - if err != nil { - logger.Errorf("Failed to listen on uapi socket: %v", err) - os.Exit(ExitSetupFailed) - } - - errs := make(chan error) - term := make(chan os.Signal, 1) - - go func() { - for { - conn, err := uapi.Accept() - if err != nil { - errs <- err - return - } - go device.IpcHandle(conn) - } - }() - logger.Verbosef("UAPI listener started") - - // wait for program to terminate - - signal.Notify(term, os.Interrupt) - signal.Notify(term, os.Kill) - signal.Notify(term, windows.SIGTERM) - - select { - case <-term: - case <-errs: - case <-device.Wait(): - } - - // clean up - - uapi.Close() - device.Close() - - logger.Verbosef("Shutting down") -} diff --git a/ratelimiter/ratelimiter_test.go b/ratelimiter/ratelimiter_test.go deleted file mode 100644 index 71140da..0000000 --- a/ratelimiter/ratelimiter_test.go +++ /dev/null @@ -1,119 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package ratelimiter - -import ( - "net/netip" - "testing" - "time" -) - -type result struct { - allowed bool - text string - wait time.Duration -} - -func TestRatelimiter(t *testing.T) { - var rate Ratelimiter - var expectedResults []result - - nano := func(nano int64) time.Duration { - return time.Nanosecond * time.Duration(nano) - } - - add := func(res result) { - expectedResults = append( - expectedResults, - res, - ) - } - - for i := 0; i < packetsBurstable; i++ { - add(result{ - allowed: true, - text: "initial burst", - }) - } - - add(result{ - allowed: false, - text: "after burst", - }) - - add(result{ - allowed: true, - wait: nano(time.Second.Nanoseconds() / packetsPerSecond), - text: "filling tokens for single packet", - }) - - add(result{ - allowed: false, - text: "not having refilled enough", - }) - - add(result{ - allowed: true, - wait: 2 * (nano(time.Second.Nanoseconds() / packetsPerSecond)), - text: "filling tokens for two packet burst", - }) - - add(result{ - allowed: true, - text: "second packet in 2 packet burst", - }) - - add(result{ - allowed: false, - text: "packet following 2 packet burst", - }) - - ips := []netip.Addr{ - netip.MustParseAddr("127.0.0.1"), - netip.MustParseAddr("192.168.1.1"), - netip.MustParseAddr("172.167.2.3"), - netip.MustParseAddr("97.231.252.215"), - netip.MustParseAddr("248.97.91.167"), - netip.MustParseAddr("188.208.233.47"), - netip.MustParseAddr("104.2.183.179"), - netip.MustParseAddr("72.129.46.120"), - netip.MustParseAddr("2001:0db8:0a0b:12f0:0000:0000:0000:0001"), - netip.MustParseAddr("f5c2:818f:c052:655a:9860:b136:6894:25f0"), - netip.MustParseAddr("b2d7:15ab:48a7:b07c:a541:f144:a9fe:54fc"), - netip.MustParseAddr("a47b:786e:1671:a22b:d6f9:4ab0:abc7:c918"), - netip.MustParseAddr("ea1e:d155:7f7a:98fb:2bf5:9483:80f6:5445"), - netip.MustParseAddr("3f0e:54a2:f5b4:cd19:a21d:58e1:3746:84c4"), - } - - now := time.Now() - rate.timeNow = func() time.Time { - return now - } - defer func() { - // Lock to avoid data race with cleanup goroutine from Init. - rate.mu.Lock() - defer rate.mu.Unlock() - - rate.timeNow = time.Now - }() - timeSleep := func(d time.Duration) { - now = now.Add(d + 1) - rate.cleanup() - } - - rate.Init() - defer rate.Close() - - for i, res := range expectedResults { - timeSleep(res.wait) - for _, ip := range ips { - allowed := rate.Allow(ip) - if allowed != res.allowed { - t.Fatalf("%d: %s: rate.Allow(%q)=%v, want %v", i, res.text, ip, allowed, res.allowed) - } - } - } -} diff --git a/replay/replay_test.go b/replay/replay_test.go deleted file mode 100644 index 8378ec3..0000000 --- a/replay/replay_test.go +++ /dev/null @@ -1,119 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package replay - -import ( - "testing" -) - -/* Ported from the linux kernel implementation - * - * - */ - -const RejectAfterMessages = 1<<64 - 1<<13 - 1 - -func TestReplay(t *testing.T) { - var filter Filter - - const T_LIM = windowSize + 1 - - testNumber := 0 - T := func(n uint64, expected bool) { - testNumber++ - if filter.ValidateCounter(n, RejectAfterMessages) != expected { - t.Fatal("Test", testNumber, "failed", n, expected) - } - } - - filter.Reset() - - T(0, true) /* 1 */ - T(1, true) /* 2 */ - T(1, false) /* 3 */ - T(9, true) /* 4 */ - T(8, true) /* 5 */ - T(7, true) /* 6 */ - T(7, false) /* 7 */ - T(T_LIM, true) /* 8 */ - T(T_LIM-1, true) /* 9 */ - T(T_LIM-1, false) /* 10 */ - T(T_LIM-2, true) /* 11 */ - T(2, true) /* 12 */ - T(2, false) /* 13 */ - T(T_LIM+16, true) /* 14 */ - T(3, false) /* 15 */ - T(T_LIM+16, false) /* 16 */ - T(T_LIM*4, true) /* 17 */ - T(T_LIM*4-(T_LIM-1), true) /* 18 */ - T(10, false) /* 19 */ - T(T_LIM*4-T_LIM, false) /* 20 */ - T(T_LIM*4-(T_LIM+1), false) /* 21 */ - T(T_LIM*4-(T_LIM-2), true) /* 22 */ - T(T_LIM*4+1-T_LIM, false) /* 23 */ - T(0, false) /* 24 */ - T(RejectAfterMessages, false) /* 25 */ - T(RejectAfterMessages-1, true) /* 26 */ - T(RejectAfterMessages, false) /* 27 */ - T(RejectAfterMessages-1, false) /* 28 */ - T(RejectAfterMessages-2, true) /* 29 */ - T(RejectAfterMessages+1, false) /* 30 */ - T(RejectAfterMessages+2, false) /* 31 */ - T(RejectAfterMessages-2, false) /* 32 */ - T(RejectAfterMessages-3, true) /* 33 */ - T(0, false) /* 34 */ - - t.Log("Bulk test 1") - filter.Reset() - testNumber = 0 - for i := uint64(1); i <= windowSize; i++ { - T(i, true) - } - T(0, true) - T(0, false) - - t.Log("Bulk test 2") - filter.Reset() - testNumber = 0 - for i := uint64(2); i <= windowSize+1; i++ { - T(i, true) - } - T(1, true) - T(0, false) - - t.Log("Bulk test 3") - filter.Reset() - testNumber = 0 - for i := uint64(windowSize + 1); i > 0; i-- { - T(i, true) - } - - t.Log("Bulk test 4") - filter.Reset() - testNumber = 0 - for i := uint64(windowSize + 2); i > 1; i-- { - T(i, true) - } - T(0, false) - - t.Log("Bulk test 5") - filter.Reset() - testNumber = 0 - for i := uint64(windowSize); i > 0; i-- { - T(i, true) - } - T(windowSize+1, true) - T(0, false) - - t.Log("Bulk test 6") - filter.Reset() - testNumber = 0 - for i := uint64(windowSize); i > 0; i-- { - T(i, true) - } - T(0, true) - T(windowSize+1, true) -} diff --git a/tai64n/tai64n_test.go b/tai64n/tai64n_test.go deleted file mode 100644 index d0b4425..0000000 --- a/tai64n/tai64n_test.go +++ /dev/null @@ -1,40 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package tai64n - -import ( - "testing" - "time" -) - -// Test that timestamps are monotonic as required by Wireguard and that -// nanosecond-level information is whitened to prevent side channel attacks. -func TestMonotonic(t *testing.T) { - startTime := time.Unix(0, 123456789) // a nontrivial bit pattern - // Whitening should reduce timestamp granularity - // to more than 10 but fewer than 20 milliseconds. - tests := []struct { - name string - t1, t2 time.Time - wantAfter bool - }{ - {"after_10_ns", startTime, startTime.Add(10 * time.Nanosecond), false}, - {"after_10_us", startTime, startTime.Add(10 * time.Microsecond), false}, - {"after_1_ms", startTime, startTime.Add(time.Millisecond), false}, - {"after_10_ms", startTime, startTime.Add(10 * time.Millisecond), false}, - {"after_20_ms", startTime, startTime.Add(20 * time.Millisecond), true}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ts1, ts2 := stamp(tt.t1), stamp(tt.t2) - got := ts2.After(ts1) - if got != tt.wantAfter { - t.Errorf("after = %v; want %v", got, tt.wantAfter) - } - }) - } -} diff --git a/tests/netns.sh b/tests/netns.sh deleted file mode 100755 index 2f2a2cd..0000000 --- a/tests/netns.sh +++ /dev/null @@ -1,425 +0,0 @@ -#!/bin/bash - -# Copyright (C) 2015-2017 Jason A. Donenfeld . All Rights Reserved. - -# This script tests the below topology: -# -# ┌─────────────────────┐ ┌──────────────────────────────────┐ ┌─────────────────────┐ -# │ $ns1 namespace │ │ $ns0 namespace │ │ $ns2 namespace │ -# │ │ │ │ │ │ -# │┌────────┐ │ │ ┌────────┐ │ │ ┌────────┐│ -# ││ wg1 │───────────┼───┼────────────│ lo │────────────┼───┼───────────│ wg2 ││ -# │├────────┴──────────┐│ │ ┌───────┴────────┴────────┐ │ │┌──────────┴────────┤│ -# ││192.168.241.1/24 ││ │ │(ns1) (ns2) │ │ ││192.168.241.2/24 ││ -# ││fd00::1/24 ││ │ │127.0.0.1:1 127.0.0.1:2│ │ ││fd00::2/24 ││ -# │└───────────────────┘│ │ │[::]:1 [::]:2 │ │ │└───────────────────┘│ -# └─────────────────────┘ │ └─────────────────────────┘ │ └─────────────────────┘ -# └──────────────────────────────────┘ -# -# After the topology is prepared we run a series of TCP/UDP iperf3 tests between the -# wireguard peers in $ns1 and $ns2. Note that $ns0 is the endpoint for the wg1 -# interfaces in $ns1 and $ns2. See https://www.wireguard.com/netns/ for further -# details on how this is accomplished. - -# This code is ported to the WireGuard-Go directly from the kernel project. -# -# Please ensure that you have installed the newest version of the WireGuard -# tools from the WireGuard project and before running these tests as: -# -# ./netns.sh - -set -e - -exec 3>&1 -export WG_HIDE_KEYS=never -netns0="wg-test-$$-0" -netns1="wg-test-$$-1" -netns2="wg-test-$$-2" -program=$1 -export LOG_LEVEL="verbose" - -pretty() { echo -e "\x1b[32m\x1b[1m[+] ${1:+NS$1: }${2}\x1b[0m" >&3; } -pp() { pretty "" "$*"; "$@"; } -maybe_exec() { if [[ $BASHPID -eq $$ ]]; then "$@"; else exec "$@"; fi; } -n0() { pretty 0 "$*"; maybe_exec ip netns exec $netns0 "$@"; } -n1() { pretty 1 "$*"; maybe_exec ip netns exec $netns1 "$@"; } -n2() { pretty 2 "$*"; maybe_exec ip netns exec $netns2 "$@"; } -ip0() { pretty 0 "ip $*"; ip -n $netns0 "$@"; } -ip1() { pretty 1 "ip $*"; ip -n $netns1 "$@"; } -ip2() { pretty 2 "ip $*"; ip -n $netns2 "$@"; } -sleep() { read -t "$1" -N 0 || true; } -waitiperf() { pretty "${1//*-}" "wait for iperf:5201"; while [[ $(ss -N "$1" -tlp 'sport = 5201') != *iperf3* ]]; do sleep 0.1; done; } -waitncatudp() { pretty "${1//*-}" "wait for udp:1111"; while [[ $(ss -N "$1" -ulp 'sport = 1111') != *ncat* ]]; do sleep 0.1; done; } -waitiface() { pretty "${1//*-}" "wait for $2 to come up"; ip netns exec "$1" bash -c "while [[ \$(< \"/sys/class/net/$2/operstate\") != up ]]; do read -t .1 -N 0 || true; done;"; } - -cleanup() { - set +e - exec 2>/dev/null - printf "$orig_message_cost" > /proc/sys/net/core/message_cost - ip0 link del dev wg1 - ip1 link del dev wg1 - ip2 link del dev wg1 - local to_kill="$(ip netns pids $netns0) $(ip netns pids $netns1) $(ip netns pids $netns2)" - [[ -n $to_kill ]] && kill $to_kill - pp ip netns del $netns1 - pp ip netns del $netns2 - pp ip netns del $netns0 - exit -} - -orig_message_cost="$(< /proc/sys/net/core/message_cost)" -trap cleanup EXIT -printf 0 > /proc/sys/net/core/message_cost - -ip netns del $netns0 2>/dev/null || true -ip netns del $netns1 2>/dev/null || true -ip netns del $netns2 2>/dev/null || true -pp ip netns add $netns0 -pp ip netns add $netns1 -pp ip netns add $netns2 -ip0 link set up dev lo - -# ip0 link add dev wg1 type wireguard -n0 $program wg1 -ip0 link set wg1 netns $netns1 - -# ip0 link add dev wg1 type wireguard -n0 $program wg2 -ip0 link set wg2 netns $netns2 - -key1="$(pp wg genkey)" -key2="$(pp wg genkey)" -pub1="$(pp wg pubkey <<<"$key1")" -pub2="$(pp wg pubkey <<<"$key2")" -psk="$(pp wg genpsk)" -[[ -n $key1 && -n $key2 && -n $psk ]] - -configure_peers() { - - ip1 addr add 192.168.241.1/24 dev wg1 - ip1 addr add fd00::1/24 dev wg1 - - ip2 addr add 192.168.241.2/24 dev wg2 - ip2 addr add fd00::2/24 dev wg2 - - n0 wg set wg1 \ - private-key <(echo "$key1") \ - listen-port 10000 \ - peer "$pub2" \ - preshared-key <(echo "$psk") \ - allowed-ips 192.168.241.2/32,fd00::2/128 - n0 wg set wg2 \ - private-key <(echo "$key2") \ - listen-port 20000 \ - peer "$pub1" \ - preshared-key <(echo "$psk") \ - allowed-ips 192.168.241.1/32,fd00::1/128 - - n0 wg showconf wg1 - n0 wg showconf wg2 - - ip1 link set up dev wg1 - ip2 link set up dev wg2 - sleep 1 -} -configure_peers - -tests() { - # Ping over IPv4 - n2 ping -c 10 -f -W 1 192.168.241.1 - n1 ping -c 10 -f -W 1 192.168.241.2 - - # Ping over IPv6 - n2 ping6 -c 10 -f -W 1 fd00::1 - n1 ping6 -c 10 -f -W 1 fd00::2 - - # TCP over IPv4 - n2 iperf3 -s -1 -B 192.168.241.2 & - waitiperf $netns2 - n1 iperf3 -Z -n 1G -c 192.168.241.2 - - # TCP over IPv6 - n1 iperf3 -s -1 -B fd00::1 & - waitiperf $netns1 - n2 iperf3 -Z -n 1G -c fd00::1 - - # UDP over IPv4 - n1 iperf3 -s -1 -B 192.168.241.1 & - waitiperf $netns1 - n2 iperf3 -Z -n 1G -b 0 -u -c 192.168.241.1 - - # UDP over IPv6 - n2 iperf3 -s -1 -B fd00::2 & - waitiperf $netns2 - n1 iperf3 -Z -n 1G -b 0 -u -c fd00::2 -} - -[[ $(ip1 link show dev wg1) =~ mtu\ ([0-9]+) ]] && orig_mtu="${BASH_REMATCH[1]}" -big_mtu=$(( 34816 - 1500 + $orig_mtu )) - -# Test using IPv4 as outer transport -n0 wg set wg1 peer "$pub2" endpoint 127.0.0.1:20000 -n0 wg set wg2 peer "$pub1" endpoint 127.0.0.1:10000 - -# Before calling tests, we first make sure that the stats counters are working -n2 ping -c 10 -f -W 1 192.168.241.1 -{ read _; read _; read _; read rx_bytes _; read _; read tx_bytes _; } < <(ip2 -stats link show dev wg2) -ip2 -stats link show dev wg2 -n0 wg show -[[ $rx_bytes -ge 840 && $tx_bytes -ge 880 && $rx_bytes -lt 2500 && $rx_bytes -lt 2500 ]] -echo "counters working" -tests -ip1 link set wg1 mtu $big_mtu -ip2 link set wg2 mtu $big_mtu -tests - -ip1 link set wg1 mtu $orig_mtu -ip2 link set wg2 mtu $orig_mtu - -# Test using IPv6 as outer transport -n0 wg set wg1 peer "$pub2" endpoint [::1]:20000 -n0 wg set wg2 peer "$pub1" endpoint [::1]:10000 -tests -ip1 link set wg1 mtu $big_mtu -ip2 link set wg2 mtu $big_mtu -tests - -ip1 link set wg1 mtu $orig_mtu -ip2 link set wg2 mtu $orig_mtu - -# Test using IPv4 that roaming works -ip0 -4 addr del 127.0.0.1/8 dev lo -ip0 -4 addr add 127.212.121.99/8 dev lo -n0 wg set wg1 listen-port 9999 -n0 wg set wg1 peer "$pub2" endpoint 127.0.0.1:20000 -n1 ping6 -W 1 -c 1 fd00::2 -[[ $(n2 wg show wg2 endpoints) == "$pub1 127.212.121.99:9999" ]] - -# Test using IPv6 that roaming works -n1 wg set wg1 listen-port 9998 -n1 wg set wg1 peer "$pub2" endpoint [::1]:20000 -n1 ping -W 1 -c 1 192.168.241.2 -[[ $(n2 wg show wg2 endpoints) == "$pub1 [::1]:9998" ]] - -# Test that crypto-RP filter works -n1 wg set wg1 peer "$pub2" allowed-ips 192.168.241.0/24 -exec 4< <(n1 ncat -l -u -p 1111) -nmap_pid=$! -waitncatudp $netns1 -n2 ncat -u 192.168.241.1 1111 <<<"X" -read -r -N 1 -t 1 out <&4 && [[ $out == "X" ]] -kill $nmap_pid -more_specific_key="$(pp wg genkey | pp wg pubkey)" -n0 wg set wg1 peer "$more_specific_key" allowed-ips 192.168.241.2/32 -n0 wg set wg2 listen-port 9997 -exec 4< <(n1 ncat -l -u -p 1111) -nmap_pid=$! -waitncatudp $netns1 -n2 ncat -u 192.168.241.1 1111 <<<"X" -! read -r -N 1 -t 1 out <&4 -kill $nmap_pid -n0 wg set wg1 peer "$more_specific_key" remove -[[ $(n1 wg show wg1 endpoints) == "$pub2 [::1]:9997" ]] - -ip1 link del wg1 -ip2 link del wg2 - -# Test using NAT. We now change the topology to this: -# ┌────────────────────────────────────────┐ ┌────────────────────────────────────────────────┐ ┌────────────────────────────────────────┐ -# │ $ns1 namespace │ │ $ns0 namespace │ │ $ns2 namespace │ -# │ │ │ │ │ │ -# │ ┌─────┐ ┌─────┐ │ │ ┌──────┐ ┌──────┐ │ │ ┌─────┐ ┌─────┐ │ -# │ │ wg1 │─────────────│vethc│───────────┼────┼────│vethrc│ │vethrs│──────────────┼─────┼──│veths│────────────│ wg2 │ │ -# │ ├─────┴──────────┐ ├─────┴──────────┐│ │ ├──────┴─────────┐ ├──────┴────────────┐ │ │ ├─────┴──────────┐ ├─────┴──────────┐ │ -# │ │192.168.241.1/24│ │192.168.1.100/24││ │ │192.168.1.100/24│ │10.0.0.1/24 │ │ │ │10.0.0.100/24 │ │192.168.241.2/24│ │ -# │ │fd00::1/24 │ │ ││ │ │ │ │SNAT:192.168.1.0/24│ │ │ │ │ │fd00::2/24 │ │ -# │ └────────────────┘ └────────────────┘│ │ └────────────────┘ └───────────────────┘ │ │ └────────────────┘ └────────────────┘ │ -# └────────────────────────────────────────┘ └────────────────────────────────────────────────┘ └────────────────────────────────────────┘ - -# ip1 link add dev wg1 type wireguard -# ip2 link add dev wg1 type wireguard - -n1 $program wg1 -n2 $program wg2 - -configure_peers - -ip0 link add vethrc type veth peer name vethc -ip0 link add vethrs type veth peer name veths -ip0 link set vethc netns $netns1 -ip0 link set veths netns $netns2 -ip0 link set vethrc up -ip0 link set vethrs up -ip0 addr add 192.168.1.1/24 dev vethrc -ip0 addr add 10.0.0.1/24 dev vethrs -ip1 addr add 192.168.1.100/24 dev vethc -ip1 link set vethc up -ip1 route add default via 192.168.1.1 -ip2 addr add 10.0.0.100/24 dev veths -ip2 link set veths up -waitiface $netns0 vethrc -waitiface $netns0 vethrs -waitiface $netns1 vethc -waitiface $netns2 veths - -n0 bash -c 'printf 1 > /proc/sys/net/ipv4/ip_forward' -n0 bash -c 'printf 2 > /proc/sys/net/netfilter/nf_conntrack_udp_timeout' -n0 bash -c 'printf 2 > /proc/sys/net/netfilter/nf_conntrack_udp_timeout_stream' -n0 iptables -t nat -A POSTROUTING -s 192.168.1.0/24 -d 10.0.0.0/24 -j SNAT --to 10.0.0.1 - -n0 wg set wg1 peer "$pub2" endpoint 10.0.0.100:20000 persistent-keepalive 1 -n1 ping -W 1 -c 1 192.168.241.2 -n2 ping -W 1 -c 1 192.168.241.1 -[[ $(n2 wg show wg2 endpoints) == "$pub1 10.0.0.1:10000" ]] -# Demonstrate n2 can still send packets to n1, since persistent-keepalive will prevent connection tracking entry from expiring (to see entries: `n0 conntrack -L`). -pp sleep 3 -n2 ping -W 1 -c 1 192.168.241.1 - -n0 iptables -t nat -F -ip0 link del vethrc -ip0 link del vethrs -ip1 link del wg1 -ip2 link del wg2 - -# Test that saddr routing is sticky but not too sticky, changing to this topology: -# ┌────────────────────────────────────────┐ ┌────────────────────────────────────────┐ -# │ $ns1 namespace │ │ $ns2 namespace │ -# │ │ │ │ -# │ ┌─────┐ ┌─────┐ │ │ ┌─────┐ ┌─────┐ │ -# │ │ wg1 │─────────────│veth1│───────────┼────┼──│veth2│────────────│ wg2 │ │ -# │ ├─────┴──────────┐ ├─────┴──────────┐│ │ ├─────┴──────────┐ ├─────┴──────────┐ │ -# │ │192.168.241.1/24│ │10.0.0.1/24 ││ │ │10.0.0.2/24 │ │192.168.241.2/24│ │ -# │ │fd00::1/24 │ │fd00:aa::1/96 ││ │ │fd00:aa::2/96 │ │fd00::2/24 │ │ -# │ └────────────────┘ └────────────────┘│ │ └────────────────┘ └────────────────┘ │ -# └────────────────────────────────────────┘ └────────────────────────────────────────┘ - -# ip1 link add dev wg1 type wireguard -# ip2 link add dev wg1 type wireguard -n1 $program wg1 -n2 $program wg2 - -configure_peers - -ip1 link add veth1 type veth peer name veth2 -ip1 link set veth2 netns $netns2 -n1 bash -c 'printf 0 > /proc/sys/net/ipv6/conf/veth1/accept_dad' -n2 bash -c 'printf 0 > /proc/sys/net/ipv6/conf/veth2/accept_dad' -n1 bash -c 'printf 1 > /proc/sys/net/ipv4/conf/veth1/promote_secondaries' - -# First we check that we aren't overly sticky and can fall over to new IPs when old ones are removed -ip1 addr add 10.0.0.1/24 dev veth1 -ip1 addr add fd00:aa::1/96 dev veth1 -ip2 addr add 10.0.0.2/24 dev veth2 -ip2 addr add fd00:aa::2/96 dev veth2 -ip1 link set veth1 up -ip2 link set veth2 up -waitiface $netns1 veth1 -waitiface $netns2 veth2 -n0 wg set wg1 peer "$pub2" endpoint 10.0.0.2:20000 -n1 ping -W 1 -c 1 192.168.241.2 -ip1 addr add 10.0.0.10/24 dev veth1 -ip1 addr del 10.0.0.1/24 dev veth1 -n1 ping -W 1 -c 1 192.168.241.2 -n0 wg set wg1 peer "$pub2" endpoint [fd00:aa::2]:20000 -n1 ping -W 1 -c 1 192.168.241.2 -ip1 addr add fd00:aa::10/96 dev veth1 -ip1 addr del fd00:aa::1/96 dev veth1 -n1 ping -W 1 -c 1 192.168.241.2 - -# Now we show that we can successfully do reply to sender routing -ip1 link set veth1 down -ip2 link set veth2 down -ip1 addr flush dev veth1 -ip2 addr flush dev veth2 -ip1 addr add 10.0.0.1/24 dev veth1 -ip1 addr add 10.0.0.2/24 dev veth1 -ip1 addr add fd00:aa::1/96 dev veth1 -ip1 addr add fd00:aa::2/96 dev veth1 -ip2 addr add 10.0.0.3/24 dev veth2 -ip2 addr add fd00:aa::3/96 dev veth2 -ip1 link set veth1 up -ip2 link set veth2 up -waitiface $netns1 veth1 -waitiface $netns2 veth2 -n0 wg set wg2 peer "$pub1" endpoint 10.0.0.1:10000 -n2 ping -W 1 -c 1 192.168.241.1 -[[ $(n0 wg show wg2 endpoints) == "$pub1 10.0.0.1:10000" ]] -n0 wg set wg2 peer "$pub1" endpoint [fd00:aa::1]:10000 -n2 ping -W 1 -c 1 192.168.241.1 -[[ $(n0 wg show wg2 endpoints) == "$pub1 [fd00:aa::1]:10000" ]] -n0 wg set wg2 peer "$pub1" endpoint 10.0.0.2:10000 -n2 ping -W 1 -c 1 192.168.241.1 -[[ $(n0 wg show wg2 endpoints) == "$pub1 10.0.0.2:10000" ]] -n0 wg set wg2 peer "$pub1" endpoint [fd00:aa::2]:10000 -n2 ping -W 1 -c 1 192.168.241.1 -[[ $(n0 wg show wg2 endpoints) == "$pub1 [fd00:aa::2]:10000" ]] - -ip1 link del veth1 -ip1 link del wg1 -ip2 link del wg2 - -# Test that Netlink/IPC is working properly by doing things that usually cause split responses - -n0 $program wg0 -sleep 5 -config=( "[Interface]" "PrivateKey=$(wg genkey)" "[Peer]" "PublicKey=$(wg genkey)" ) -for a in {1..255}; do - for b in {0..255}; do - config+=( "AllowedIPs=$a.$b.0.0/16,$a::$b/128" ) - done -done -n0 wg setconf wg0 <(printf '%s\n' "${config[@]}") -i=0 -for ip in $(n0 wg show wg0 allowed-ips); do - ((++i)) -done -((i == 255*256*2+1)) -ip0 link del wg0 - -n0 $program wg0 -config=( "[Interface]" "PrivateKey=$(wg genkey)" ) -for a in {1..40}; do - config+=( "[Peer]" "PublicKey=$(wg genkey)" ) - for b in {1..52}; do - config+=( "AllowedIPs=$a.$b.0.0/16" ) - done -done -n0 wg setconf wg0 <(printf '%s\n' "${config[@]}") -i=0 -while read -r line; do - j=0 - for ip in $line; do - ((++j)) - done - ((j == 53)) - ((++i)) -done < <(n0 wg show wg0 allowed-ips) -((i == 40)) -ip0 link del wg0 - -n0 $program wg0 -config=( ) -for i in {1..29}; do - config+=( "[Peer]" "PublicKey=$(wg genkey)" ) -done -config+=( "[Peer]" "PublicKey=$(wg genkey)" "AllowedIPs=255.2.3.4/32,abcd::255/128" ) -n0 wg setconf wg0 <(printf '%s\n' "${config[@]}") -n0 wg showconf wg0 > /dev/null -ip0 link del wg0 - -! n0 wg show doesnotexist || false - -declare -A objects -while read -t 0.1 -r line 2>/dev/null || [[ $? -ne 142 ]]; do - [[ $line =~ .*(wg[0-9]+:\ [A-Z][a-z]+\ [0-9]+)\ .*(created|destroyed).* ]] || continue - objects["${BASH_REMATCH[1]}"]+="${BASH_REMATCH[2]}" -done < /dev/kmsg -alldeleted=1 -for object in "${!objects[@]}"; do - if [[ ${objects["$object"]} != *createddestroyed ]]; then - echo "Error: $object: merely ${objects["$object"]}" >&3 - alldeleted=0 - fi -done -[[ $alldeleted -eq 1 ]] -pretty "" "Objects that were created were also destroyed." diff --git a/tun/alignment_windows_test.go b/tun/alignment_windows_test.go deleted file mode 100644 index e3252b2..0000000 --- a/tun/alignment_windows_test.go +++ /dev/null @@ -1,67 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package tun - -import ( - "reflect" - "testing" - "unsafe" -) - -func checkAlignment(t *testing.T, name string, offset uintptr) { - t.Helper() - if offset%8 != 0 { - t.Errorf("offset of %q within struct is %d bytes, which does not align to 64-bit word boundaries (missing %d bytes). Atomic operations will crash on 32-bit systems.", name, offset, 8-(offset%8)) - } -} - -// TestRateJugglerAlignment checks that atomically-accessed fields are -// aligned to 64-bit boundaries, as required by the atomic package. -// -// Unfortunately, violating this rule on 32-bit platforms results in a -// hard segfault at runtime. -func TestRateJugglerAlignment(t *testing.T) { - var r rateJuggler - - typ := reflect.TypeOf(&r).Elem() - t.Logf("Peer type size: %d, with fields:", typ.Size()) - for i := 0; i < typ.NumField(); i++ { - field := typ.Field(i) - t.Logf("\t%30s\toffset=%3v\t(type size=%3d, align=%d)", - field.Name, - field.Offset, - field.Type.Size(), - field.Type.Align(), - ) - } - - checkAlignment(t, "rateJuggler.current", unsafe.Offsetof(r.current)) - checkAlignment(t, "rateJuggler.nextByteCount", unsafe.Offsetof(r.nextByteCount)) - checkAlignment(t, "rateJuggler.nextStartTime", unsafe.Offsetof(r.nextStartTime)) -} - -// TestNativeTunAlignment checks that atomically-accessed fields are -// aligned to 64-bit boundaries, as required by the atomic package. -// -// Unfortunately, violating this rule on 32-bit platforms results in a -// hard segfault at runtime. -func TestNativeTunAlignment(t *testing.T) { - var tun NativeTun - - typ := reflect.TypeOf(&tun).Elem() - t.Logf("Peer type size: %d, with fields:", typ.Size()) - for i := 0; i < typ.NumField(); i++ { - field := typ.Field(i) - t.Logf("\t%30s\toffset=%3v\t(type size=%3d, align=%d)", - field.Name, - field.Offset, - field.Type.Size(), - field.Type.Align(), - ) - } - - checkAlignment(t, "NativeTun.rate", unsafe.Offsetof(tun.rate)) -} diff --git a/tun/checksum_test.go b/tun/checksum_test.go deleted file mode 100644 index 4ea9b8b..0000000 --- a/tun/checksum_test.go +++ /dev/null @@ -1,98 +0,0 @@ -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, - 128, - 256, - 512, - 1024, - 1500, - 2048, - 4096, - 8192, - 9000, - 9001, - } - - for _, length := range lengths { - b.Run(fmt.Sprintf("%d", length), func(b *testing.B) { - buf := make([]byte, length) - rng := rand.New(rand.NewSource(1)) - rng.Read(buf) - b.ResetTimer() - for i := 0; i < b.N; i++ { - checksum(buf, 0) - } - }) - } -} diff --git a/tun/netstack/examples/http_client.go b/tun/netstack/examples/http_client.go deleted file mode 100644 index d71267d..0000000 --- a/tun/netstack/examples/http_client.go +++ /dev/null @@ -1,54 +0,0 @@ -//go:build ignore - -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package main - -import ( - "io" - "log" - "net/http" - "net/netip" - - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/device" - "golang.zx2c4.com/wireguard/tun/netstack" -) - -func main() { - tun, tnet, err := netstack.CreateNetTUN( - []netip.Addr{netip.MustParseAddr("192.168.4.28")}, - []netip.Addr{netip.MustParseAddr("8.8.8.8")}, - 1420) - if err != nil { - log.Panic(err) - } - dev := device.NewDevice(tun, conn.NewDefaultBind(), device.NewLogger(device.LogLevelVerbose, "")) - err = dev.IpcSet(`private_key=087ec6e14bbed210e7215cdc73468dfa23f080a1bfb8665b2fd809bd99d28379 -public_key=c4c8e984c5322c8184c72265b92b250fdb63688705f504ba003c88f03393cf28 -allowed_ip=0.0.0.0/0 -endpoint=127.0.0.1:58120 -`) - err = dev.Up() - if err != nil { - log.Panic(err) - } - - client := http.Client{ - Transport: &http.Transport{ - DialContext: tnet.DialContext, - }, - } - resp, err := client.Get("http://192.168.4.29/") - if err != nil { - log.Panic(err) - } - body, err := io.ReadAll(resp.Body) - if err != nil { - log.Panic(err) - } - log.Println(string(body)) -} diff --git a/tun/netstack/examples/http_server.go b/tun/netstack/examples/http_server.go deleted file mode 100644 index 7278851..0000000 --- a/tun/netstack/examples/http_server.go +++ /dev/null @@ -1,51 +0,0 @@ -//go:build ignore - -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package main - -import ( - "io" - "log" - "net" - "net/http" - "net/netip" - - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/device" - "golang.zx2c4.com/wireguard/tun/netstack" -) - -func main() { - tun, tnet, err := netstack.CreateNetTUN( - []netip.Addr{netip.MustParseAddr("192.168.4.29")}, - []netip.Addr{netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("8.8.4.4")}, - 1420, - ) - if err != nil { - log.Panic(err) - } - dev := device.NewDevice(tun, conn.NewDefaultBind(), device.NewLogger(device.LogLevelVerbose, "")) - dev.IpcSet(`private_key=003ed5d73b55806c30de3f8a7bdab38af13539220533055e635690b8b87ad641 -listen_port=58120 -public_key=f928d4f6c1b86c12f2562c10b07c555c5c57fd00f59e90c8d8d88767271cbf7c -allowed_ip=192.168.4.28/32 -persistent_keepalive_interval=25 -`) - dev.Up() - listener, err := tnet.ListenTCP(&net.TCPAddr{Port: 80}) - if err != nil { - log.Panicln(err) - } - http.HandleFunc("/", func(writer http.ResponseWriter, request *http.Request) { - log.Printf("> %s - %s - %s", request.RemoteAddr, request.URL.String(), request.UserAgent()) - io.WriteString(writer, "Hello from userspace TCP!") - }) - err = http.Serve(listener, nil) - if err != nil { - log.Panicln(err) - } -} diff --git a/tun/netstack/examples/ping_client.go b/tun/netstack/examples/ping_client.go deleted file mode 100644 index d1b562f..0000000 --- a/tun/netstack/examples/ping_client.go +++ /dev/null @@ -1,75 +0,0 @@ -//go:build ignore - -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package main - -import ( - "bytes" - "log" - "math/rand" - "net/netip" - "time" - - "golang.org/x/net/icmp" - "golang.org/x/net/ipv4" - - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/device" - "golang.zx2c4.com/wireguard/tun/netstack" -) - -func main() { - tun, tnet, err := netstack.CreateNetTUN( - []netip.Addr{netip.MustParseAddr("192.168.4.29")}, - []netip.Addr{netip.MustParseAddr("8.8.8.8")}, - 1420) - if err != nil { - log.Panic(err) - } - dev := device.NewDevice(tun, conn.NewDefaultBind(), device.NewLogger(device.LogLevelVerbose, "")) - dev.IpcSet(`private_key=a8dac1d8a70a751f0f699fb14ba1cff7b79cf4fbd8f09f44c6e6a90d0369604f -public_key=25123c5dcd3328ff645e4f2a3fce0d754400d3887a0cb7c56f0267e20fbf3c5b -endpoint=163.172.161.0:12912 -allowed_ip=0.0.0.0/0 -`) - err = dev.Up() - if err != nil { - log.Panic(err) - } - - socket, err := tnet.Dial("ping4", "zx2c4.com") - if err != nil { - log.Panic(err) - } - requestPing := icmp.Echo{ - Seq: rand.Intn(1 << 16), - Data: []byte("gopher burrow"), - } - icmpBytes, _ := (&icmp.Message{Type: ipv4.ICMPTypeEcho, Code: 0, Body: &requestPing}).Marshal(nil) - socket.SetReadDeadline(time.Now().Add(time.Second * 10)) - start := time.Now() - _, err = socket.Write(icmpBytes) - if err != nil { - log.Panic(err) - } - n, err := socket.Read(icmpBytes[:]) - if err != nil { - log.Panic(err) - } - replyPacket, err := icmp.ParseMessage(1, icmpBytes[:n]) - if err != nil { - log.Panic(err) - } - replyPing, ok := replyPacket.Body.(*icmp.Echo) - if !ok { - log.Panicf("invalid reply type: %v", replyPacket) - } - if !bytes.Equal(replyPing.Data, requestPing.Data) || replyPing.Seq != requestPing.Seq { - log.Panicf("invalid ping reply: %v", replyPing) - } - log.Printf("Ping latency: %v", time.Since(start)) -} diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go deleted file mode 100644 index a7aec9e..0000000 --- a/tun/netstack/tun.go +++ /dev/null @@ -1,1057 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package netstack - -import ( - "bytes" - "context" - "crypto/rand" - "encoding/binary" - "errors" - "fmt" - "io" - "net" - "net/netip" - "os" - "regexp" - "strconv" - "strings" - "syscall" - "time" - - "golang.zx2c4.com/wireguard/tun" - - "golang.org/x/net/dns/dnsmessage" - "gvisor.dev/gvisor/pkg/buffer" - "gvisor.dev/gvisor/pkg/tcpip" - "gvisor.dev/gvisor/pkg/tcpip/adapters/gonet" - "gvisor.dev/gvisor/pkg/tcpip/header" - "gvisor.dev/gvisor/pkg/tcpip/link/channel" - "gvisor.dev/gvisor/pkg/tcpip/network/ipv4" - "gvisor.dev/gvisor/pkg/tcpip/network/ipv6" - "gvisor.dev/gvisor/pkg/tcpip/stack" - "gvisor.dev/gvisor/pkg/tcpip/transport/icmp" - "gvisor.dev/gvisor/pkg/tcpip/transport/tcp" - "gvisor.dev/gvisor/pkg/tcpip/transport/udp" - "gvisor.dev/gvisor/pkg/waiter" -) - -type netTun struct { - ep *channel.Endpoint - stack *stack.Stack - events chan tun.Event - notifyHandle *channel.NotificationHandle - incomingPacket chan *buffer.View - mtu int - dnsServers []netip.Addr - hasV4, hasV6 bool -} - -type Net netTun - -func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net, error) { - opts := stack.Options{ - NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, - TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}, - HandleLocal: true, - } - dev := &netTun{ - ep: channel.New(1024, uint32(mtu), ""), - stack: stack.New(opts), - events: make(chan tun.Event, 10), - incomingPacket: make(chan *buffer.View), - dnsServers: dnsServers, - mtu: mtu, - } - sackEnabledOpt := tcpip.TCPSACKEnabled(true) // TCP SACK is disabled by default - tcpipErr := dev.stack.SetTransportProtocolOption(tcp.ProtocolNumber, &sackEnabledOpt) - if tcpipErr != nil { - return nil, nil, fmt.Errorf("could not enable TCP SACK: %v", tcpipErr) - } - dev.notifyHandle = dev.ep.AddNotify(dev) - tcpipErr = dev.stack.CreateNIC(1, dev.ep) - if tcpipErr != nil { - return nil, nil, fmt.Errorf("CreateNIC: %v", tcpipErr) - } - for _, ip := range localAddresses { - var protoNumber tcpip.NetworkProtocolNumber - if ip.Is4() { - protoNumber = ipv4.ProtocolNumber - } else if ip.Is6() { - protoNumber = ipv6.ProtocolNumber - } - protoAddr := tcpip.ProtocolAddress{ - Protocol: protoNumber, - AddressWithPrefix: tcpip.AddrFromSlice(ip.AsSlice()).WithPrefix(), - } - tcpipErr := dev.stack.AddProtocolAddress(1, protoAddr, stack.AddressProperties{}) - if tcpipErr != nil { - return nil, nil, fmt.Errorf("AddProtocolAddress(%v): %v", ip, tcpipErr) - } - if ip.Is4() { - dev.hasV4 = true - } else if ip.Is6() { - dev.hasV6 = true - } - } - if dev.hasV4 { - dev.stack.AddRoute(tcpip.Route{Destination: header.IPv4EmptySubnet, NIC: 1}) - } - if dev.hasV6 { - dev.stack.AddRoute(tcpip.Route{Destination: header.IPv6EmptySubnet, NIC: 1}) - } - - dev.events <- tun.EventUp - return dev, (*Net)(dev), nil -} - -func (tun *netTun) Name() (string, error) { - return "go", nil -} - -func (tun *netTun) File() *os.File { - return nil -} - -func (tun *netTun) Events() <-chan tun.Event { - return tun.events -} - -func (tun *netTun) Read(buf [][]byte, sizes []int, offset int) (int, error) { - view, ok := <-tun.incomingPacket - if !ok { - return 0, os.ErrClosed - } - - n, err := view.Read(buf[0][offset:]) - if err != nil { - return 0, err - } - sizes[0] = n - return 1, nil -} - -func (tun *netTun) Write(buf [][]byte, offset int) (int, error) { - for _, buf := range buf { - packet := buf[offset:] - if len(packet) == 0 { - continue - } - - pkb := stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buffer.MakeWithData(packet)}) - switch packet[0] >> 4 { - case 4: - tun.ep.InjectInbound(header.IPv4ProtocolNumber, pkb) - case 6: - tun.ep.InjectInbound(header.IPv6ProtocolNumber, pkb) - default: - return 0, syscall.EAFNOSUPPORT - } - } - return len(buf), nil -} - -func (tun *netTun) WriteNotify() { - pkt := tun.ep.Read() - if pkt == nil { - return - } - - view := pkt.ToView() - pkt.DecRef() - - tun.incomingPacket <- view -} - -func (tun *netTun) Close() error { - tun.stack.RemoveNIC(1) - tun.stack.Close() - tun.ep.RemoveNotify(tun.notifyHandle) - tun.ep.Close() - - if tun.events != nil { - close(tun.events) - } - - if tun.incomingPacket != nil { - close(tun.incomingPacket) - } - - return nil -} - -func (tun *netTun) MTU() (int, error) { - return tun.mtu, nil -} - -func (tun *netTun) BatchSize() int { - return 1 -} - -func convertToFullAddr(endpoint netip.AddrPort) (tcpip.FullAddress, tcpip.NetworkProtocolNumber) { - var protoNumber tcpip.NetworkProtocolNumber - if endpoint.Addr().Is4() { - protoNumber = ipv4.ProtocolNumber - } else { - protoNumber = ipv6.ProtocolNumber - } - return tcpip.FullAddress{ - NIC: 1, - Addr: tcpip.AddrFromSlice(endpoint.Addr().AsSlice()), - Port: endpoint.Port(), - }, protoNumber -} - -func (net *Net) DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (*gonet.TCPConn, error) { - fa, pn := convertToFullAddr(addr) - return gonet.DialContextTCP(ctx, net.stack, fa, pn) -} - -func (net *Net) DialContextTCP(ctx context.Context, addr *net.TCPAddr) (*gonet.TCPConn, error) { - if addr == nil { - return net.DialContextTCPAddrPort(ctx, netip.AddrPort{}) - } - ip, _ := netip.AddrFromSlice(addr.IP) - return net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(addr.Port))) -} - -func (net *Net) DialTCPAddrPort(addr netip.AddrPort) (*gonet.TCPConn, error) { - fa, pn := convertToFullAddr(addr) - return gonet.DialTCP(net.stack, fa, pn) -} - -func (net *Net) DialTCP(addr *net.TCPAddr) (*gonet.TCPConn, error) { - if addr == nil { - return net.DialTCPAddrPort(netip.AddrPort{}) - } - ip, _ := netip.AddrFromSlice(addr.IP) - return net.DialTCPAddrPort(netip.AddrPortFrom(ip, uint16(addr.Port))) -} - -func (net *Net) ListenTCPAddrPort(addr netip.AddrPort) (*gonet.TCPListener, error) { - fa, pn := convertToFullAddr(addr) - return gonet.ListenTCP(net.stack, fa, pn) -} - -func (net *Net) ListenTCP(addr *net.TCPAddr) (*gonet.TCPListener, error) { - if addr == nil { - return net.ListenTCPAddrPort(netip.AddrPort{}) - } - ip, _ := netip.AddrFromSlice(addr.IP) - return net.ListenTCPAddrPort(netip.AddrPortFrom(ip, uint16(addr.Port))) -} - -func (net *Net) DialUDPAddrPort(laddr, raddr netip.AddrPort) (*gonet.UDPConn, error) { - var lfa, rfa *tcpip.FullAddress - var pn tcpip.NetworkProtocolNumber - if laddr.IsValid() || laddr.Port() > 0 { - var addr tcpip.FullAddress - addr, pn = convertToFullAddr(laddr) - lfa = &addr - } - if raddr.IsValid() || raddr.Port() > 0 { - var addr tcpip.FullAddress - addr, pn = convertToFullAddr(raddr) - rfa = &addr - } - return gonet.DialUDP(net.stack, lfa, rfa, pn) -} - -func (net *Net) ListenUDPAddrPort(laddr netip.AddrPort) (*gonet.UDPConn, error) { - return net.DialUDPAddrPort(laddr, netip.AddrPort{}) -} - -func (net *Net) DialUDP(laddr, raddr *net.UDPAddr) (*gonet.UDPConn, error) { - var la, ra netip.AddrPort - if laddr != nil { - ip, _ := netip.AddrFromSlice(laddr.IP) - la = netip.AddrPortFrom(ip, uint16(laddr.Port)) - } - if raddr != nil { - ip, _ := netip.AddrFromSlice(raddr.IP) - ra = netip.AddrPortFrom(ip, uint16(raddr.Port)) - } - return net.DialUDPAddrPort(la, ra) -} - -func (net *Net) ListenUDP(laddr *net.UDPAddr) (*gonet.UDPConn, error) { - return net.DialUDP(laddr, nil) -} - -type PingConn struct { - laddr PingAddr - raddr PingAddr - wq waiter.Queue - ep tcpip.Endpoint - deadline *time.Timer -} - -type PingAddr struct{ addr netip.Addr } - -func (ia PingAddr) String() string { - return ia.addr.String() -} - -func (ia PingAddr) Network() string { - if ia.addr.Is4() { - return "ping4" - } else if ia.addr.Is6() { - return "ping6" - } - return "ping" -} - -func (ia PingAddr) Addr() netip.Addr { - return ia.addr -} - -func PingAddrFromAddr(addr netip.Addr) *PingAddr { - return &PingAddr{addr} -} - -func (net *Net) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { - if !laddr.IsValid() && !raddr.IsValid() { - return nil, errors.New("ping dial: invalid address") - } - v6 := laddr.Is6() || raddr.Is6() - bind := laddr.IsValid() - if !bind { - if v6 { - laddr = netip.IPv6Unspecified() - } else { - laddr = netip.IPv4Unspecified() - } - } - - tn := icmp.ProtocolNumber4 - pn := ipv4.ProtocolNumber - if v6 { - tn = icmp.ProtocolNumber6 - pn = ipv6.ProtocolNumber - } - - pc := &PingConn{ - laddr: PingAddr{laddr}, - deadline: time.NewTimer(time.Hour << 10), - } - pc.deadline.Stop() - - ep, tcpipErr := net.stack.NewEndpoint(tn, pn, &pc.wq) - if tcpipErr != nil { - return nil, fmt.Errorf("ping socket: endpoint: %s", tcpipErr) - } - pc.ep = ep - - if bind { - fa, _ := convertToFullAddr(netip.AddrPortFrom(laddr, 0)) - if tcpipErr = pc.ep.Bind(fa); tcpipErr != nil { - return nil, fmt.Errorf("ping bind: %s", tcpipErr) - } - } - - if raddr.IsValid() { - pc.raddr = PingAddr{raddr} - fa, _ := convertToFullAddr(netip.AddrPortFrom(raddr, 0)) - if tcpipErr = pc.ep.Connect(fa); tcpipErr != nil { - return nil, fmt.Errorf("ping connect: %s", tcpipErr) - } - } - - return pc, nil -} - -func (net *Net) ListenPingAddr(laddr netip.Addr) (*PingConn, error) { - return net.DialPingAddr(laddr, netip.Addr{}) -} - -func (net *Net) DialPing(laddr, raddr *PingAddr) (*PingConn, error) { - var la, ra netip.Addr - if laddr != nil { - la = laddr.addr - } - if raddr != nil { - ra = raddr.addr - } - return net.DialPingAddr(la, ra) -} - -func (net *Net) ListenPing(laddr *PingAddr) (*PingConn, error) { - var la netip.Addr - if laddr != nil { - la = laddr.addr - } - return net.ListenPingAddr(la) -} - -func (pc *PingConn) LocalAddr() net.Addr { - return pc.laddr -} - -func (pc *PingConn) RemoteAddr() net.Addr { - return pc.raddr -} - -func (pc *PingConn) Close() error { - pc.deadline.Reset(0) - pc.ep.Close() - return nil -} - -func (pc *PingConn) SetWriteDeadline(t time.Time) error { - return errors.New("not implemented") -} - -func (pc *PingConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { - var na netip.Addr - switch v := addr.(type) { - case *PingAddr: - na = v.addr - case *net.IPAddr: - na, _ = netip.AddrFromSlice(v.IP) - default: - return 0, fmt.Errorf("ping write: wrong net.Addr type") - } - if !((na.Is4() && pc.laddr.addr.Is4()) || (na.Is6() && pc.laddr.addr.Is6())) { - return 0, fmt.Errorf("ping write: mismatched protocols") - } - - buf := bytes.NewReader(p) - rfa, _ := convertToFullAddr(netip.AddrPortFrom(na, 0)) - // won't block, no deadlines - n64, tcpipErr := pc.ep.Write(buf, tcpip.WriteOptions{ - To: &rfa, - }) - if tcpipErr != nil { - return int(n64), fmt.Errorf("ping write: %s", tcpipErr) - } - - return int(n64), nil -} - -func (pc *PingConn) Write(p []byte) (n int, err error) { - return pc.WriteTo(p, &pc.raddr) -} - -func (pc *PingConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { - e, notifyCh := waiter.NewChannelEntry(waiter.EventIn) - pc.wq.EventRegister(&e) - defer pc.wq.EventUnregister(&e) - - select { - case <-pc.deadline.C: - return 0, nil, os.ErrDeadlineExceeded - case <-notifyCh: - } - - w := tcpip.SliceWriter(p) - - res, tcpipErr := pc.ep.Read(&w, tcpip.ReadOptions{ - NeedRemoteAddr: true, - }) - if tcpipErr != nil { - return 0, nil, fmt.Errorf("ping read: %s", tcpipErr) - } - - remoteAddr, _ := netip.AddrFromSlice(res.RemoteAddr.Addr.AsSlice()) - return res.Count, &PingAddr{remoteAddr}, nil -} - -func (pc *PingConn) Read(p []byte) (n int, err error) { - n, _, err = pc.ReadFrom(p) - return -} - -func (pc *PingConn) SetDeadline(t time.Time) error { - // pc.SetWriteDeadline is unimplemented - - return pc.SetReadDeadline(t) -} - -func (pc *PingConn) SetReadDeadline(t time.Time) error { - pc.deadline.Reset(time.Until(t)) - return nil -} - -var ( - errNoSuchHost = errors.New("no such host") - errLameReferral = errors.New("lame referral") - errCannotUnmarshalDNSMessage = errors.New("cannot unmarshal DNS message") - errCannotMarshalDNSMessage = errors.New("cannot marshal DNS message") - errServerMisbehaving = errors.New("server misbehaving") - errInvalidDNSResponse = errors.New("invalid DNS response") - errNoAnswerFromDNSServer = errors.New("no answer from DNS server") - errServerTemporarilyMisbehaving = errors.New("server misbehaving") - errCanceled = errors.New("operation was canceled") - errTimeout = errors.New("i/o timeout") - errNumericPort = errors.New("port must be numeric") - errNoSuitableAddress = errors.New("no suitable address found") - errMissingAddress = errors.New("missing address") -) - -func (net *Net) LookupHost(host string) (addrs []string, err error) { - return net.LookupContextHost(context.Background(), host) -} - -func isDomainName(s string) bool { - l := len(s) - if l == 0 || l > 254 || l == 254 && s[l-1] != '.' { - return false - } - last := byte('.') - nonNumeric := false - partlen := 0 - for i := 0; i < len(s); i++ { - c := s[i] - switch { - default: - return false - case 'a' <= c && c <= 'z' || 'A' <= c && c <= 'Z' || c == '_': - nonNumeric = true - partlen++ - case '0' <= c && c <= '9': - partlen++ - case c == '-': - if last == '.' { - return false - } - partlen++ - nonNumeric = true - case c == '.': - if last == '.' || last == '-' { - return false - } - if partlen > 63 || partlen == 0 { - return false - } - partlen = 0 - } - last = c - } - if last == '-' || partlen > 63 { - return false - } - return nonNumeric -} - -func randU16() uint16 { - var b [2]byte - _, err := rand.Read(b[:]) - if err != nil { - panic(err) - } - return binary.LittleEndian.Uint16(b[:]) -} - -func newRequest(q dnsmessage.Question) (id uint16, udpReq, tcpReq []byte, err error) { - id = randU16() - b := dnsmessage.NewBuilder(make([]byte, 2, 514), dnsmessage.Header{ID: id, RecursionDesired: true}) - b.EnableCompression() - if err := b.StartQuestions(); err != nil { - return 0, nil, nil, err - } - if err := b.Question(q); err != nil { - return 0, nil, nil, err - } - tcpReq, err = b.Finish() - udpReq = tcpReq[2:] - l := len(tcpReq) - 2 - tcpReq[0] = byte(l >> 8) - tcpReq[1] = byte(l) - return id, udpReq, tcpReq, err -} - -func equalASCIIName(x, y dnsmessage.Name) bool { - if x.Length != y.Length { - return false - } - for i := 0; i < int(x.Length); i++ { - a := x.Data[i] - b := y.Data[i] - if 'A' <= a && a <= 'Z' { - a += 0x20 - } - if 'A' <= b && b <= 'Z' { - b += 0x20 - } - if a != b { - return false - } - } - return true -} - -func checkResponse(reqID uint16, reqQues dnsmessage.Question, respHdr dnsmessage.Header, respQues dnsmessage.Question) bool { - if !respHdr.Response { - return false - } - if reqID != respHdr.ID { - return false - } - if reqQues.Type != respQues.Type || reqQues.Class != respQues.Class || !equalASCIIName(reqQues.Name, respQues.Name) { - return false - } - return true -} - -func dnsPacketRoundTrip(c net.Conn, id uint16, query dnsmessage.Question, b []byte) (dnsmessage.Parser, dnsmessage.Header, error) { - if _, err := c.Write(b); err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, err - } - b = make([]byte, 512) - for { - n, err := c.Read(b) - if err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, err - } - var p dnsmessage.Parser - h, err := p.Start(b[:n]) - if err != nil { - continue - } - q, err := p.Question() - if err != nil || !checkResponse(id, query, h, q) { - continue - } - return p, h, nil - } -} - -func dnsStreamRoundTrip(c net.Conn, id uint16, query dnsmessage.Question, b []byte) (dnsmessage.Parser, dnsmessage.Header, error) { - if _, err := c.Write(b); err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, err - } - b = make([]byte, 1280) - if _, err := io.ReadFull(c, b[:2]); err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, err - } - l := int(b[0])<<8 | int(b[1]) - if l > len(b) { - b = make([]byte, l) - } - n, err := io.ReadFull(c, b[:l]) - if err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, err - } - var p dnsmessage.Parser - h, err := p.Start(b[:n]) - if err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, errCannotUnmarshalDNSMessage - } - q, err := p.Question() - if err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, errCannotUnmarshalDNSMessage - } - if !checkResponse(id, query, h, q) { - return dnsmessage.Parser{}, dnsmessage.Header{}, errInvalidDNSResponse - } - return p, h, nil -} - -func (tnet *Net) exchange(ctx context.Context, server netip.Addr, q dnsmessage.Question, timeout time.Duration) (dnsmessage.Parser, dnsmessage.Header, error) { - q.Class = dnsmessage.ClassINET - id, udpReq, tcpReq, err := newRequest(q) - if err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, errCannotMarshalDNSMessage - } - - for _, useUDP := range []bool{true, false} { - ctx, cancel := context.WithDeadline(ctx, time.Now().Add(timeout)) - defer cancel() - - var c net.Conn - var err error - if useUDP { - c, err = tnet.DialUDPAddrPort(netip.AddrPort{}, netip.AddrPortFrom(server, 53)) - } else { - c, err = tnet.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(server, 53)) - } - - if err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, err - } - if d, ok := ctx.Deadline(); ok && !d.IsZero() { - err := c.SetDeadline(d) - if err != nil { - return dnsmessage.Parser{}, dnsmessage.Header{}, err - } - } - var p dnsmessage.Parser - var h dnsmessage.Header - if useUDP { - p, h, err = dnsPacketRoundTrip(c, id, q, udpReq) - } else { - p, h, err = dnsStreamRoundTrip(c, id, q, tcpReq) - } - c.Close() - if err != nil { - if err == context.Canceled { - err = errCanceled - } else if err == context.DeadlineExceeded { - err = errTimeout - } - return dnsmessage.Parser{}, dnsmessage.Header{}, err - } - if err := p.SkipQuestion(); err != dnsmessage.ErrSectionDone { - return dnsmessage.Parser{}, dnsmessage.Header{}, errInvalidDNSResponse - } - if h.Truncated { - continue - } - return p, h, nil - } - return dnsmessage.Parser{}, dnsmessage.Header{}, errNoAnswerFromDNSServer -} - -func checkHeader(p *dnsmessage.Parser, h dnsmessage.Header) error { - if h.RCode == dnsmessage.RCodeNameError { - return errNoSuchHost - } - _, err := p.AnswerHeader() - if err != nil && err != dnsmessage.ErrSectionDone { - return errCannotUnmarshalDNSMessage - } - if h.RCode == dnsmessage.RCodeSuccess && !h.Authoritative && !h.RecursionAvailable && err == dnsmessage.ErrSectionDone { - return errLameReferral - } - if h.RCode != dnsmessage.RCodeSuccess && h.RCode != dnsmessage.RCodeNameError { - if h.RCode == dnsmessage.RCodeServerFailure { - return errServerTemporarilyMisbehaving - } - return errServerMisbehaving - } - return nil -} - -func skipToAnswer(p *dnsmessage.Parser, qtype dnsmessage.Type) error { - for { - h, err := p.AnswerHeader() - if err == dnsmessage.ErrSectionDone { - return errNoSuchHost - } - if err != nil { - return errCannotUnmarshalDNSMessage - } - if h.Type == qtype { - return nil - } - if err := p.SkipAnswer(); err != nil { - return errCannotUnmarshalDNSMessage - } - } -} - -func (tnet *Net) tryOneName(ctx context.Context, name string, qtype dnsmessage.Type) (dnsmessage.Parser, string, error) { - var lastErr error - - n, err := dnsmessage.NewName(name) - if err != nil { - return dnsmessage.Parser{}, "", errCannotMarshalDNSMessage - } - q := dnsmessage.Question{ - Name: n, - Type: qtype, - Class: dnsmessage.ClassINET, - } - - for i := 0; i < 2; i++ { - for _, server := range tnet.dnsServers { - p, h, err := tnet.exchange(ctx, server, q, time.Second*5) - if err != nil { - dnsErr := &net.DNSError{ - Err: err.Error(), - Name: name, - Server: server.String(), - } - if nerr, ok := err.(net.Error); ok && nerr.Timeout() { - dnsErr.IsTimeout = true - } - if _, ok := err.(*net.OpError); ok { - dnsErr.IsTemporary = true - } - lastErr = dnsErr - continue - } - - if err := checkHeader(&p, h); err != nil { - dnsErr := &net.DNSError{ - Err: err.Error(), - Name: name, - Server: server.String(), - } - if err == errServerTemporarilyMisbehaving { - dnsErr.IsTemporary = true - } - if err == errNoSuchHost { - dnsErr.IsNotFound = true - return p, server.String(), dnsErr - } - lastErr = dnsErr - continue - } - - err = skipToAnswer(&p, qtype) - if err == nil { - return p, server.String(), nil - } - lastErr = &net.DNSError{ - Err: err.Error(), - Name: name, - Server: server.String(), - } - if err == errNoSuchHost { - lastErr.(*net.DNSError).IsNotFound = true - return p, server.String(), lastErr - } - } - } - return dnsmessage.Parser{}, "", lastErr -} - -func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, error) { - if host == "" || (!tnet.hasV6 && !tnet.hasV4) { - return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} - } - zlen := len(host) - if strings.IndexByte(host, ':') != -1 { - if zidx := strings.LastIndexByte(host, '%'); zidx != -1 { - zlen = zidx - } - } - if ip, err := netip.ParseAddr(host[:zlen]); err == nil { - return []string{ip.String()}, nil - } - - if !isDomainName(host) { - return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} - } - type result struct { - p dnsmessage.Parser - server string - error - } - var addrsV4, addrsV6 []netip.Addr - lanes := 0 - if tnet.hasV4 { - lanes++ - } - if tnet.hasV6 { - lanes++ - } - lane := make(chan result, lanes) - var lastErr error - if tnet.hasV4 { - go func() { - p, server, err := tnet.tryOneName(ctx, host+".", dnsmessage.TypeA) - lane <- result{p, server, err} - }() - } - if tnet.hasV6 { - go func() { - p, server, err := tnet.tryOneName(ctx, host+".", dnsmessage.TypeAAAA) - lane <- result{p, server, err} - }() - } - for l := 0; l < lanes; l++ { - result := <-lane - if result.error != nil { - if lastErr == nil { - lastErr = result.error - } - continue - } - - loop: - for { - h, err := result.p.AnswerHeader() - if err != nil && err != dnsmessage.ErrSectionDone { - lastErr = &net.DNSError{ - Err: errCannotMarshalDNSMessage.Error(), - Name: host, - Server: result.server, - } - } - if err != nil { - break - } - switch h.Type { - case dnsmessage.TypeA: - a, err := result.p.AResource() - if err != nil { - lastErr = &net.DNSError{ - Err: errCannotMarshalDNSMessage.Error(), - Name: host, - Server: result.server, - } - break loop - } - addrsV4 = append(addrsV4, netip.AddrFrom4(a.A)) - - case dnsmessage.TypeAAAA: - aaaa, err := result.p.AAAAResource() - if err != nil { - lastErr = &net.DNSError{ - Err: errCannotMarshalDNSMessage.Error(), - Name: host, - Server: result.server, - } - break loop - } - addrsV6 = append(addrsV6, netip.AddrFrom16(aaaa.AAAA)) - - default: - if err := result.p.SkipAnswer(); err != nil { - lastErr = &net.DNSError{ - Err: errCannotMarshalDNSMessage.Error(), - Name: host, - Server: result.server, - } - break loop - } - continue - } - } - } - // We don't do RFC6724. Instead just put V6 addresses first if an IPv6 address is enabled - var addrs []netip.Addr - if tnet.hasV6 { - addrs = append(addrsV6, addrsV4...) - } else { - addrs = append(addrsV4, addrsV6...) - } - - if len(addrs) == 0 && lastErr != nil { - return nil, lastErr - } - saddrs := make([]string, 0, len(addrs)) - for _, ip := range addrs { - saddrs = append(saddrs, ip.String()) - } - return saddrs, nil -} - -func partialDeadline(now, deadline time.Time, addrsRemaining int) (time.Time, error) { - if deadline.IsZero() { - return deadline, nil - } - timeRemaining := deadline.Sub(now) - if timeRemaining <= 0 { - return time.Time{}, errTimeout - } - timeout := timeRemaining / time.Duration(addrsRemaining) - const saneMinimum = 2 * time.Second - if timeout < saneMinimum { - if timeRemaining < saneMinimum { - timeout = timeRemaining - } else { - timeout = saneMinimum - } - } - return now.Add(timeout), nil -} - -var protoSplitter = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`) - -func (tnet *Net) DialContext(ctx context.Context, network, address string) (net.Conn, error) { - if ctx == nil { - panic("nil context") - } - var acceptV4, acceptV6 bool - matches := protoSplitter.FindStringSubmatch(network) - if matches == nil { - return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)} - } else if len(matches[2]) == 0 { - acceptV4 = true - acceptV6 = true - } else { - acceptV4 = matches[2][0] == '4' - acceptV6 = !acceptV4 - } - var host string - var port int - if matches[1] == "ping" { - host = address - } else { - var sport string - var err error - host, sport, err = net.SplitHostPort(address) - if err != nil { - return nil, &net.OpError{Op: "dial", Err: err} - } - port, err = strconv.Atoi(sport) - if err != nil || port < 0 || port > 65535 { - return nil, &net.OpError{Op: "dial", Err: errNumericPort} - } - } - allAddr, err := tnet.LookupContextHost(ctx, host) - if err != nil { - return nil, &net.OpError{Op: "dial", Err: err} - } - var addrs []netip.AddrPort - for _, addr := range allAddr { - ip, err := netip.ParseAddr(addr) - if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) { - addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port))) - } - } - if len(addrs) == 0 && len(allAddr) != 0 { - return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress} - } - - var firstErr error - for i, addr := range addrs { - select { - case <-ctx.Done(): - err := ctx.Err() - if err == context.Canceled { - err = errCanceled - } else if err == context.DeadlineExceeded { - err = errTimeout - } - return nil, &net.OpError{Op: "dial", Err: err} - default: - } - - dialCtx := ctx - if deadline, hasDeadline := ctx.Deadline(); hasDeadline { - partialDeadline, err := partialDeadline(time.Now(), deadline, len(addrs)-i) - if err != nil { - if firstErr == nil { - firstErr = &net.OpError{Op: "dial", Err: err} - } - break - } - if partialDeadline.Before(deadline) { - var cancel context.CancelFunc - dialCtx, cancel = context.WithDeadline(ctx, partialDeadline) - defer cancel() - } - } - - var c net.Conn - switch matches[1] { - case "tcp": - c, err = tnet.DialContextTCPAddrPort(dialCtx, addr) - case "udp": - c, err = tnet.DialUDPAddrPort(netip.AddrPort{}, addr) - case "ping": - c, err = tnet.DialPingAddr(netip.Addr{}, addr.Addr()) - } - if err == nil { - return c, nil - } - if firstErr == nil { - firstErr = err - } - } - if firstErr == nil { - firstErr = &net.OpError{Op: "dial", Err: errMissingAddress} - } - return nil, firstErr -} - -func (tnet *Net) Dial(network, address string) (net.Conn, error) { - return tnet.DialContext(context.Background(), network, address) -} diff --git a/tun/offload_linux_test.go b/tun/offload_linux_test.go deleted file mode 100644 index d87e636..0000000 --- a/tun/offload_linux_test.go +++ /dev/null @@ -1,752 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package tun - -import ( - "net/netip" - "testing" - - "golang.org/x/sys/unix" - "golang.zx2c4.com/wireguard/conn" - "gvisor.dev/gvisor/pkg/tcpip" - "gvisor.dev/gvisor/pkg/tcpip/header" -) - -const ( - offset = virtioNetHdrLen -) - -var ( - ip4PortA = netip.MustParseAddrPort("192.0.2.1:1") - ip4PortB = netip.MustParseAddrPort("192.0.2.2:1") - ip4PortC = netip.MustParseAddrPort("192.0.2.3:1") - ip6PortA = netip.MustParseAddrPort("[2001:db8::1]:1") - ip6PortB = netip.MustParseAddrPort("[2001:db8::2]:1") - ip6PortC = netip.MustParseAddrPort("[2001:db8::3]:1") -) - -func udp4PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, payloadLen int, ipFn func(*header.IPv4Fields)) []byte { - totalLen := 28 + payloadLen - b := make([]byte, offset+int(totalLen), 65535) - ipv4H := header.IPv4(b[offset:]) - srcAs4 := srcIPPort.Addr().As4() - dstAs4 := dstIPPort.Addr().As4() - ipFields := &header.IPv4Fields{ - SrcAddr: tcpip.AddrFromSlice(srcAs4[:]), - DstAddr: tcpip.AddrFromSlice(dstAs4[:]), - Protocol: unix.IPPROTO_UDP, - TTL: 64, - TotalLength: uint16(totalLen), - } - if ipFn != nil { - ipFn(ipFields) - } - ipv4H.Encode(ipFields) - udpH := header.UDP(b[offset+20:]) - udpH.Encode(&header.UDPFields{ - SrcPort: srcIPPort.Port(), - DstPort: dstIPPort.Port(), - Length: uint16(payloadLen + udphLen), - }) - ipv4H.SetChecksum(^ipv4H.CalculateChecksum()) - pseudoCsum := header.PseudoHeaderChecksum(unix.IPPROTO_UDP, ipv4H.SourceAddress(), ipv4H.DestinationAddress(), uint16(udphLen+payloadLen)) - udpH.SetChecksum(^udpH.CalculateChecksum(pseudoCsum)) - return b -} - -func udp6Packet(srcIPPort, dstIPPort netip.AddrPort, payloadLen int) []byte { - return udp6PacketMutateIPFields(srcIPPort, dstIPPort, payloadLen, nil) -} - -func udp6PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, payloadLen int, ipFn func(*header.IPv6Fields)) []byte { - totalLen := 48 + payloadLen - b := make([]byte, offset+int(totalLen), 65535) - ipv6H := header.IPv6(b[offset:]) - srcAs16 := srcIPPort.Addr().As16() - dstAs16 := dstIPPort.Addr().As16() - ipFields := &header.IPv6Fields{ - SrcAddr: tcpip.AddrFromSlice(srcAs16[:]), - DstAddr: tcpip.AddrFromSlice(dstAs16[:]), - TransportProtocol: unix.IPPROTO_UDP, - HopLimit: 64, - PayloadLength: uint16(payloadLen + udphLen), - } - if ipFn != nil { - ipFn(ipFields) - } - ipv6H.Encode(ipFields) - udpH := header.UDP(b[offset+40:]) - udpH.Encode(&header.UDPFields{ - SrcPort: srcIPPort.Port(), - DstPort: dstIPPort.Port(), - Length: uint16(payloadLen + udphLen), - }) - pseudoCsum := header.PseudoHeaderChecksum(unix.IPPROTO_UDP, ipv6H.SourceAddress(), ipv6H.DestinationAddress(), uint16(udphLen+payloadLen)) - udpH.SetChecksum(^udpH.CalculateChecksum(pseudoCsum)) - return b -} - -func udp4Packet(srcIPPort, dstIPPort netip.AddrPort, payloadLen int) []byte { - return udp4PacketMutateIPFields(srcIPPort, dstIPPort, payloadLen, nil) -} - -func tcp4PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, flags header.TCPFlags, segmentSize, seq uint32, ipFn func(*header.IPv4Fields)) []byte { - totalLen := 40 + segmentSize - b := make([]byte, offset+int(totalLen), 65535) - ipv4H := header.IPv4(b[offset:]) - srcAs4 := srcIPPort.Addr().As4() - dstAs4 := dstIPPort.Addr().As4() - ipFields := &header.IPv4Fields{ - SrcAddr: tcpip.AddrFromSlice(srcAs4[:]), - DstAddr: tcpip.AddrFromSlice(dstAs4[:]), - Protocol: unix.IPPROTO_TCP, - TTL: 64, - TotalLength: uint16(totalLen), - } - if ipFn != nil { - ipFn(ipFields) - } - ipv4H.Encode(ipFields) - tcpH := header.TCP(b[offset+20:]) - tcpH.Encode(&header.TCPFields{ - SrcPort: srcIPPort.Port(), - DstPort: dstIPPort.Port(), - SeqNum: seq, - AckNum: 1, - DataOffset: 20, - Flags: flags, - WindowSize: 3000, - }) - ipv4H.SetChecksum(^ipv4H.CalculateChecksum()) - pseudoCsum := header.PseudoHeaderChecksum(unix.IPPROTO_TCP, ipv4H.SourceAddress(), ipv4H.DestinationAddress(), uint16(20+segmentSize)) - tcpH.SetChecksum(^tcpH.CalculateChecksum(pseudoCsum)) - return b -} - -func tcp4Packet(srcIPPort, dstIPPort netip.AddrPort, flags header.TCPFlags, segmentSize, seq uint32) []byte { - return tcp4PacketMutateIPFields(srcIPPort, dstIPPort, flags, segmentSize, seq, nil) -} - -func tcp6PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, flags header.TCPFlags, segmentSize, seq uint32, ipFn func(*header.IPv6Fields)) []byte { - totalLen := 60 + segmentSize - b := make([]byte, offset+int(totalLen), 65535) - ipv6H := header.IPv6(b[offset:]) - srcAs16 := srcIPPort.Addr().As16() - dstAs16 := dstIPPort.Addr().As16() - ipFields := &header.IPv6Fields{ - SrcAddr: tcpip.AddrFromSlice(srcAs16[:]), - DstAddr: tcpip.AddrFromSlice(dstAs16[:]), - TransportProtocol: unix.IPPROTO_TCP, - HopLimit: 64, - PayloadLength: uint16(segmentSize + 20), - } - if ipFn != nil { - ipFn(ipFields) - } - ipv6H.Encode(ipFields) - tcpH := header.TCP(b[offset+40:]) - tcpH.Encode(&header.TCPFields{ - SrcPort: srcIPPort.Port(), - DstPort: dstIPPort.Port(), - SeqNum: seq, - AckNum: 1, - DataOffset: 20, - Flags: flags, - WindowSize: 3000, - }) - pseudoCsum := header.PseudoHeaderChecksum(unix.IPPROTO_TCP, ipv6H.SourceAddress(), ipv6H.DestinationAddress(), uint16(20+segmentSize)) - tcpH.SetChecksum(^tcpH.CalculateChecksum(pseudoCsum)) - return b -} - -func tcp6Packet(srcIPPort, dstIPPort netip.AddrPort, flags header.TCPFlags, segmentSize, seq uint32) []byte { - return tcp6PacketMutateIPFields(srcIPPort, dstIPPort, flags, segmentSize, seq, nil) -} - -func Test_handleVirtioRead(t *testing.T) { - tests := []struct { - name string - hdr virtioNetHdr - pktIn []byte - wantLens []int - wantErr bool - }{ - { - "tcp4", - virtioNetHdr{ - flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, - gsoType: unix.VIRTIO_NET_HDR_GSO_TCPV4, - gsoSize: 100, - hdrLen: 40, - csumStart: 20, - csumOffset: 16, - }, - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck|header.TCPFlagPsh, 200, 1), - []int{140, 140}, - false, - }, - { - "tcp6", - virtioNetHdr{ - flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, - gsoType: unix.VIRTIO_NET_HDR_GSO_TCPV6, - gsoSize: 100, - hdrLen: 60, - csumStart: 40, - csumOffset: 16, - }, - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck|header.TCPFlagPsh, 200, 1), - []int{160, 160}, - false, - }, - { - "udp4", - virtioNetHdr{ - flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, - gsoType: unix.VIRTIO_NET_HDR_GSO_UDP_L4, - gsoSize: 100, - hdrLen: 28, - csumStart: 20, - csumOffset: 6, - }, - udp4Packet(ip4PortA, ip4PortB, 200), - []int{128, 128}, - false, - }, - { - "udp6", - virtioNetHdr{ - flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, - gsoType: unix.VIRTIO_NET_HDR_GSO_UDP_L4, - gsoSize: 100, - hdrLen: 48, - csumStart: 40, - csumOffset: 6, - }, - udp6Packet(ip6PortA, ip6PortB, 200), - []int{148, 148}, - false, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - out := make([][]byte, conn.IdealBatchSize) - sizes := make([]int, conn.IdealBatchSize) - for i := range out { - out[i] = make([]byte, 65535) - } - tt.hdr.encode(tt.pktIn) - n, err := handleVirtioRead(tt.pktIn, out, sizes, offset) - if err != nil { - if tt.wantErr { - return - } - t.Fatalf("got err: %v", err) - } - if n != len(tt.wantLens) { - t.Fatalf("got %d packets, wanted %d", n, len(tt.wantLens)) - } - for i := range tt.wantLens { - if tt.wantLens[i] != sizes[i] { - t.Fatalf("wantLens[%d]: %d != outSizes: %d", i, tt.wantLens[i], sizes[i]) - } - } - }) - } -} - -func flipTCP4Checksum(b []byte) []byte { - at := virtioNetHdrLen + 20 + 16 // 20 byte ipv4 header; tcp csum offset is 16 - b[at] ^= 0xFF - b[at+1] ^= 0xFF - return b -} - -func flipUDP4Checksum(b []byte) []byte { - at := virtioNetHdrLen + 20 + 6 // 20 byte ipv4 header; udp csum offset is 6 - b[at] ^= 0xFF - b[at+1] ^= 0xFF - return b -} - -func Fuzz_handleGRO(f *testing.F) { - pkt0 := tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1) - pkt1 := tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101) - pkt2 := tcp4Packet(ip4PortA, ip4PortC, header.TCPFlagAck, 100, 201) - pkt3 := tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1) - pkt4 := tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101) - pkt5 := tcp6Packet(ip6PortA, ip6PortC, header.TCPFlagAck, 100, 201) - pkt6 := udp4Packet(ip4PortA, ip4PortB, 100) - pkt7 := udp4Packet(ip4PortA, ip4PortB, 100) - pkt8 := udp4Packet(ip4PortA, ip4PortC, 100) - pkt9 := udp6Packet(ip6PortA, ip6PortB, 100) - pkt10 := udp6Packet(ip6PortA, ip6PortB, 100) - pkt11 := udp6Packet(ip6PortA, ip6PortC, 100) - f.Add(pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11, true, offset) - f.Fuzz(func(t *testing.T, pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11 []byte, canUDPGRO bool, offset int) { - pkts := [][]byte{pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11} - toWrite := make([]int, 0, len(pkts)) - handleGRO(pkts, offset, newTCPGROTable(), newUDPGROTable(), canUDPGRO, &toWrite) - if len(toWrite) > len(pkts) { - t.Errorf("len(toWrite): %d > len(pkts): %d", len(toWrite), len(pkts)) - } - seenWriteI := make(map[int]bool) - for _, writeI := range toWrite { - if writeI < 0 || writeI > len(pkts)-1 { - t.Errorf("toWrite value (%d) outside bounds of len(pkts): %d", writeI, len(pkts)) - } - if seenWriteI[writeI] { - t.Errorf("duplicate toWrite value: %d", writeI) - } - seenWriteI[writeI] = true - } - }) -} - -func Test_handleGRO(t *testing.T) { - tests := []struct { - name string - pktsIn [][]byte - canUDPGRO bool - wantToWrite []int - wantLens []int - wantErr bool - }{ - { - "multiple protocols and flows", - [][]byte{ - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // tcp4 flow 1 - udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1 - udp4Packet(ip4PortA, ip4PortC, 100), // udp4 flow 2 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // tcp4 flow 1 - tcp4Packet(ip4PortA, ip4PortC, header.TCPFlagAck, 100, 201), // tcp4 flow 2 - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), // tcp6 flow 1 - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101), // tcp6 flow 1 - tcp6Packet(ip6PortA, ip6PortC, header.TCPFlagAck, 100, 201), // tcp6 flow 2 - udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1 - udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1 - udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1 - }, - true, - []int{0, 1, 2, 4, 5, 7, 9}, - []int{240, 228, 128, 140, 260, 160, 248}, - false, - }, - { - "multiple protocols and flows no UDP GRO", - [][]byte{ - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // tcp4 flow 1 - udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1 - udp4Packet(ip4PortA, ip4PortC, 100), // udp4 flow 2 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // tcp4 flow 1 - tcp4Packet(ip4PortA, ip4PortC, header.TCPFlagAck, 100, 201), // tcp4 flow 2 - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), // tcp6 flow 1 - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101), // tcp6 flow 1 - tcp6Packet(ip6PortA, ip6PortC, header.TCPFlagAck, 100, 201), // tcp6 flow 2 - udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1 - udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1 - udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1 - }, - false, - []int{0, 1, 2, 4, 5, 7, 8, 9, 10}, - []int{240, 128, 128, 140, 260, 160, 128, 148, 148}, - false, - }, - { - "PSH interleaved", - [][]byte{ - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // v4 flow 1 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck|header.TCPFlagPsh, 100, 101), // v4 flow 1 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 201), // v4 flow 1 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 301), // v4 flow 1 - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), // v6 flow 1 - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck|header.TCPFlagPsh, 100, 101), // v6 flow 1 - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 201), // v6 flow 1 - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 301), // v6 flow 1 - }, - true, - []int{0, 2, 4, 6}, - []int{240, 240, 260, 260}, - false, - }, - { - "coalesceItemInvalidCSum", - [][]byte{ - flipTCP4Checksum(tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1)), // v4 flow 1 seq 1 len 100 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // v4 flow 1 seq 101 len 100 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 201), // v4 flow 1 seq 201 len 100 - flipUDP4Checksum(udp4Packet(ip4PortA, ip4PortB, 100)), - udp4Packet(ip4PortA, ip4PortB, 100), - udp4Packet(ip4PortA, ip4PortB, 100), - }, - true, - []int{0, 1, 3, 4}, - []int{140, 240, 128, 228}, - false, - }, - { - "out of order", - [][]byte{ - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // v4 flow 1 seq 101 len 100 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // v4 flow 1 seq 1 len 100 - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 201), // v4 flow 1 seq 201 len 100 - }, - true, - []int{0}, - []int{340}, - false, - }, - { - "unequal TTL", - [][]byte{ - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), - tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) { - fields.TTL++ - }), - udp4Packet(ip4PortA, ip4PortB, 100), - udp4PacketMutateIPFields(ip4PortA, ip4PortB, 100, func(fields *header.IPv4Fields) { - fields.TTL++ - }), - }, - true, - []int{0, 1, 2, 3}, - []int{140, 140, 128, 128}, - false, - }, - { - "unequal ToS", - [][]byte{ - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), - tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) { - fields.TOS++ - }), - udp4Packet(ip4PortA, ip4PortB, 100), - udp4PacketMutateIPFields(ip4PortA, ip4PortB, 100, func(fields *header.IPv4Fields) { - fields.TOS++ - }), - }, - true, - []int{0, 1, 2, 3}, - []int{140, 140, 128, 128}, - false, - }, - { - "unequal flags more fragments set", - [][]byte{ - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), - tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) { - fields.Flags = 1 - }), - udp4Packet(ip4PortA, ip4PortB, 100), - udp4PacketMutateIPFields(ip4PortA, ip4PortB, 100, func(fields *header.IPv4Fields) { - fields.Flags = 1 - }), - }, - true, - []int{0, 1, 2, 3}, - []int{140, 140, 128, 128}, - false, - }, - { - "unequal flags DF set", - [][]byte{ - tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), - tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) { - fields.Flags = 2 - }), - udp4Packet(ip4PortA, ip4PortB, 100), - udp4PacketMutateIPFields(ip4PortA, ip4PortB, 100, func(fields *header.IPv4Fields) { - fields.Flags = 2 - }), - }, - true, - []int{0, 1, 2, 3}, - []int{140, 140, 128, 128}, - false, - }, - { - "ipv6 unequal hop limit", - [][]byte{ - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), - tcp6PacketMutateIPFields(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv6Fields) { - fields.HopLimit++ - }), - udp6Packet(ip6PortA, ip6PortB, 100), - udp6PacketMutateIPFields(ip6PortA, ip6PortB, 100, func(fields *header.IPv6Fields) { - fields.HopLimit++ - }), - }, - true, - []int{0, 1, 2, 3}, - []int{160, 160, 148, 148}, - false, - }, - { - "ipv6 unequal traffic class", - [][]byte{ - tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), - tcp6PacketMutateIPFields(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv6Fields) { - fields.TrafficClass++ - }), - udp6Packet(ip6PortA, ip6PortB, 100), - udp6PacketMutateIPFields(ip6PortA, ip6PortB, 100, func(fields *header.IPv6Fields) { - fields.TrafficClass++ - }), - }, - true, - []int{0, 1, 2, 3}, - []int{160, 160, 148, 148}, - false, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - toWrite := make([]int, 0, len(tt.pktsIn)) - err := handleGRO(tt.pktsIn, offset, newTCPGROTable(), newUDPGROTable(), tt.canUDPGRO, &toWrite) - if err != nil { - if tt.wantErr { - return - } - t.Fatalf("got err: %v", err) - } - if len(toWrite) != len(tt.wantToWrite) { - t.Fatalf("got %d packets, wanted %d", len(toWrite), len(tt.wantToWrite)) - } - for i, pktI := range tt.wantToWrite { - if tt.wantToWrite[i] != toWrite[i] { - t.Fatalf("wantToWrite[%d]: %d != toWrite: %d", i, tt.wantToWrite[i], toWrite[i]) - } - if tt.wantLens[i] != len(tt.pktsIn[pktI][offset:]) { - t.Errorf("wanted len %d packet at %d, got: %d", tt.wantLens[i], i, len(tt.pktsIn[pktI][offset:])) - } - } - }) - } -} - -func Test_packetIsGROCandidate(t *testing.T) { - tcp4 := tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1)[virtioNetHdrLen:] - tcp4TooShort := tcp4[:39] - ip4InvalidHeaderLen := make([]byte, len(tcp4)) - copy(ip4InvalidHeaderLen, tcp4) - ip4InvalidHeaderLen[0] = 0x46 - ip4InvalidProtocol := make([]byte, len(tcp4)) - copy(ip4InvalidProtocol, tcp4) - ip4InvalidProtocol[9] = unix.IPPROTO_GRE - - tcp6 := tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1)[virtioNetHdrLen:] - tcp6TooShort := tcp6[:59] - ip6InvalidProtocol := make([]byte, len(tcp6)) - copy(ip6InvalidProtocol, tcp6) - ip6InvalidProtocol[6] = unix.IPPROTO_GRE - - udp4 := udp4Packet(ip4PortA, ip4PortB, 100)[virtioNetHdrLen:] - udp4TooShort := udp4[:27] - - udp6 := udp6Packet(ip6PortA, ip6PortB, 100)[virtioNetHdrLen:] - udp6TooShort := udp6[:47] - - tests := []struct { - name string - b []byte - canUDPGRO bool - want groCandidateType - }{ - { - "tcp4", - tcp4, - true, - tcp4GROCandidate, - }, - { - "tcp6", - tcp6, - true, - tcp6GROCandidate, - }, - { - "udp4", - udp4, - true, - udp4GROCandidate, - }, - { - "udp4 no support", - udp4, - false, - notGROCandidate, - }, - { - "udp6", - udp6, - true, - udp6GROCandidate, - }, - { - "udp6 no support", - udp6, - false, - notGROCandidate, - }, - { - "udp4 too short", - udp4TooShort, - true, - notGROCandidate, - }, - { - "udp6 too short", - udp6TooShort, - true, - notGROCandidate, - }, - { - "tcp4 too short", - tcp4TooShort, - true, - notGROCandidate, - }, - { - "tcp6 too short", - tcp6TooShort, - true, - notGROCandidate, - }, - { - "invalid IP version", - []byte{0x00}, - true, - notGROCandidate, - }, - { - "invalid IP header len", - ip4InvalidHeaderLen, - true, - notGROCandidate, - }, - { - "ip4 invalid protocol", - ip4InvalidProtocol, - true, - notGROCandidate, - }, - { - "ip6 invalid protocol", - ip6InvalidProtocol, - true, - notGROCandidate, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := packetIsGROCandidate(tt.b, tt.canUDPGRO); got != tt.want { - t.Errorf("packetIsGROCandidate() = %v, want %v", got, tt.want) - } - }) - } -} - -func Test_udpPacketsCanCoalesce(t *testing.T) { - udp4a := udp4Packet(ip4PortA, ip4PortB, 100) - udp4b := udp4Packet(ip4PortA, ip4PortB, 100) - udp4c := udp4Packet(ip4PortA, ip4PortB, 110) - - type args struct { - pkt []byte - iphLen uint8 - gsoSize uint16 - item udpGROItem - bufs [][]byte - bufsOffset int - } - tests := []struct { - name string - args args - want canCoalesce - }{ - { - "coalesceAppend equal gso", - args{ - pkt: udp4a[offset:], - iphLen: 20, - gsoSize: 100, - item: udpGROItem{ - gsoSize: 100, - iphLen: 20, - }, - bufs: [][]byte{ - udp4a, - udp4b, - }, - bufsOffset: offset, - }, - coalesceAppend, - }, - { - "coalesceAppend smaller gso", - args{ - pkt: udp4a[offset : len(udp4a)-90], - iphLen: 20, - gsoSize: 10, - item: udpGROItem{ - gsoSize: 100, - iphLen: 20, - }, - bufs: [][]byte{ - udp4a, - udp4b, - }, - bufsOffset: offset, - }, - coalesceAppend, - }, - { - "coalesceUnavailable smaller gso previously appended", - args{ - pkt: udp4a[offset:], - iphLen: 20, - gsoSize: 100, - item: udpGROItem{ - gsoSize: 100, - iphLen: 20, - }, - bufs: [][]byte{ - udp4c, - udp4b, - }, - bufsOffset: offset, - }, - coalesceUnavailable, - }, - { - "coalesceUnavailable larger following smaller", - args{ - pkt: udp4c[offset:], - iphLen: 20, - gsoSize: 110, - item: udpGROItem{ - gsoSize: 100, - iphLen: 20, - }, - bufs: [][]byte{ - udp4a, - udp4c, - }, - bufsOffset: offset, - }, - coalesceUnavailable, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := udpPacketsCanCoalesce(tt.args.pkt, tt.args.iphLen, tt.args.gsoSize, tt.args.item, tt.args.bufs, tt.args.bufsOffset); got != tt.want { - t.Errorf("udpPacketsCanCoalesce() = %v, want %v", got, tt.want) - } - }) - } -} diff --git a/tun/tuntest/tuntest.go b/tun/tuntest/tuntest.go deleted file mode 100644 index 9c4564f..0000000 --- a/tun/tuntest/tuntest.go +++ /dev/null @@ -1,155 +0,0 @@ -/* SPDX-License-Identifier: MIT - * - * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. - */ - -package tuntest - -import ( - "encoding/binary" - "io" - "net/netip" - "os" - - "golang.zx2c4.com/wireguard/tun" -) - -func Ping(dst, src netip.Addr) []byte { - localPort := uint16(1337) - seq := uint16(0) - - payload := make([]byte, 4) - binary.BigEndian.PutUint16(payload[0:], localPort) - binary.BigEndian.PutUint16(payload[2:], seq) - - return genICMPv4(payload, dst, src) -} - -// Checksum is the "internet checksum" from https://tools.ietf.org/html/rfc1071. -func checksum(buf []byte, initial uint16) uint16 { - v := uint32(initial) - for i := 0; i < len(buf)-1; i += 2 { - v += uint32(binary.BigEndian.Uint16(buf[i:])) - } - if len(buf)%2 == 1 { - v += uint32(buf[len(buf)-1]) << 8 - } - for v > 0xffff { - v = (v >> 16) + (v & 0xffff) - } - return ^uint16(v) -} - -func genICMPv4(payload []byte, dst, src netip.Addr) []byte { - const ( - icmpv4ProtocolNumber = 1 - icmpv4Echo = 8 - icmpv4ChecksumOffset = 2 - icmpv4Size = 8 - ipv4Size = 20 - ipv4TotalLenOffset = 2 - ipv4ChecksumOffset = 10 - ttl = 65 - headerSize = ipv4Size + icmpv4Size - ) - - pkt := make([]byte, headerSize+len(payload)) - - ip := pkt[0:ipv4Size] - icmpv4 := pkt[ipv4Size : ipv4Size+icmpv4Size] - - // https://tools.ietf.org/html/rfc792 - icmpv4[0] = icmpv4Echo // type - icmpv4[1] = 0 // code - chksum := ^checksum(icmpv4, checksum(payload, 0)) - binary.BigEndian.PutUint16(icmpv4[icmpv4ChecksumOffset:], chksum) - - // https://tools.ietf.org/html/rfc760 section 3.1 - length := uint16(len(pkt)) - ip[0] = (4 << 4) | (ipv4Size / 4) - binary.BigEndian.PutUint16(ip[ipv4TotalLenOffset:], length) - ip[8] = ttl - ip[9] = icmpv4ProtocolNumber - copy(ip[12:], src.AsSlice()) - copy(ip[16:], dst.AsSlice()) - chksum = ^checksum(ip[:], 0) - binary.BigEndian.PutUint16(ip[ipv4ChecksumOffset:], chksum) - - copy(pkt[headerSize:], payload) - return pkt -} - -type ChannelTUN struct { - Inbound chan []byte // incoming packets, closed on TUN close - Outbound chan []byte // outbound packets, blocks forever on TUN close - - closed chan struct{} - events chan tun.Event - tun chTun -} - -func NewChannelTUN() *ChannelTUN { - c := &ChannelTUN{ - Inbound: make(chan []byte), - Outbound: make(chan []byte), - closed: make(chan struct{}), - events: make(chan tun.Event, 1), - } - c.tun.c = c - c.events <- tun.EventUp - return c -} - -func (c *ChannelTUN) TUN() tun.Device { - return &c.tun -} - -type chTun struct { - c *ChannelTUN -} - -func (t *chTun) File() *os.File { return nil } - -func (t *chTun) Read(packets [][]byte, sizes []int, offset int) (int, error) { - select { - case <-t.c.closed: - return 0, os.ErrClosed - case msg := <-t.c.Outbound: - n := copy(packets[0][offset:], msg) - sizes[0] = n - return 1, nil - } -} - -// Write is called by the wireguard device to deliver a packet for routing. -func (t *chTun) Write(packets [][]byte, offset int) (int, error) { - if offset == -1 { - close(t.c.closed) - close(t.c.events) - return 0, io.EOF - } - for i, data := range packets { - msg := make([]byte, len(data)-offset) - copy(msg, data[offset:]) - select { - case <-t.c.closed: - return i, os.ErrClosed - case t.c.Inbound <- msg: - } - } - return len(packets), nil -} - -func (t *chTun) BatchSize() int { - return 1 -} - -const DefaultMTU = 1420 - -func (t *chTun) MTU() (int, error) { return DefaultMTU, nil } -func (t *chTun) Name() (string, error) { return "loopbackTun1", nil } -func (t *chTun) Events() <-chan tun.Event { return t.c.events } -func (t *chTun) Close() error { - t.Write(nil, -1) - return nil -} From 75d0f348d59d974bbbc25e74c8d1c15ce8aad28b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 17 May 2026 20:08:42 +0800 Subject: [PATCH 23/34] Rename module --- conn/bind_std.go | 10 +++------- conn/bind_windows.go | 3 +-- device/device.go | 8 ++++---- device/keypair.go | 2 +- device/noise-protocol.go | 3 +-- device/peer.go | 2 +- device/queueconstants_android.go | 2 +- device/queueconstants_default.go | 2 +- device/receive.go | 3 +-- device/send.go | 4 ++-- device/sticky_default.go | 4 ++-- device/sticky_linux.go | 5 ++--- device/tun.go | 2 +- device/uapi.go | 2 +- go.mod | 2 +- ipc/uapi_linux.go | 3 +-- ipc/uapi_windows.go | 2 +- tun/errors.go | 10 ++++------ tun/offload_linux.go | 2 +- tun/tun_linux.go | 8 +++----- 20 files changed, 33 insertions(+), 46 deletions(-) diff --git a/conn/bind_std.go b/conn/bind_std.go index f5c8816..80e8210 100644 --- a/conn/bind_std.go +++ b/conn/bind_std.go @@ -20,9 +20,7 @@ import ( "golang.org/x/net/ipv6" ) -var ( - _ Bind = (*StdNetBind)(nil) -) +var _ Bind = (*StdNetBind)(nil) // StdNetBind implements Bind for all platforms. While Windows has its own Bind // (see bind_windows.go), it may fall back to StdNetBind. @@ -210,10 +208,8 @@ func (s *StdNetBind) getMessages() *[]ipv6.Message { return s.msgsPool.Get().(*[]ipv6.Message) } -var ( - // If compilation fails here these are no longer the same underlying type. - _ ipv6.Message = ipv4.Message{} -) +// If compilation fails here these are no longer the same underlying type. +var _ ipv6.Message = ipv4.Message{} type batchReader interface { ReadBatch([]ipv6.Message, int) (int, error) diff --git a/conn/bind_windows.go b/conn/bind_windows.go index a3b8460..d166227 100644 --- a/conn/bind_windows.go +++ b/conn/bind_windows.go @@ -15,9 +15,8 @@ import ( "sync/atomic" "unsafe" + "github.com/sagernet/wireguard-go/conn/winrio" "golang.org/x/sys/windows" - - "golang.zx2c4.com/wireguard/conn/winrio" ) const ( diff --git a/device/device.go b/device/device.go index 6854ed8..bfc56bb 100644 --- a/device/device.go +++ b/device/device.go @@ -11,10 +11,10 @@ import ( "sync/atomic" "time" - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/ratelimiter" - "golang.zx2c4.com/wireguard/rwcancel" - "golang.zx2c4.com/wireguard/tun" + "github.com/sagernet/wireguard-go/conn" + "github.com/sagernet/wireguard-go/ratelimiter" + "github.com/sagernet/wireguard-go/rwcancel" + "github.com/sagernet/wireguard-go/tun" ) type Device struct { diff --git a/device/keypair.go b/device/keypair.go index 0b72e19..0704748 100644 --- a/device/keypair.go +++ b/device/keypair.go @@ -11,7 +11,7 @@ import ( "sync/atomic" "time" - "golang.zx2c4.com/wireguard/replay" + "github.com/sagernet/wireguard-go/replay" ) /* Due to limitations in Go and /x/crypto there is currently diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 5cf1702..a9b8498 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -12,11 +12,10 @@ import ( "sync" "time" + "github.com/sagernet/wireguard-go/tai64n" "golang.org/x/crypto/blake2s" "golang.org/x/crypto/chacha20poly1305" "golang.org/x/crypto/poly1305" - - "golang.zx2c4.com/wireguard/tai64n" ) type handshakeState int diff --git a/device/peer.go b/device/peer.go index ebf25f9..fff1cf5 100644 --- a/device/peer.go +++ b/device/peer.go @@ -12,7 +12,7 @@ import ( "sync/atomic" "time" - "golang.zx2c4.com/wireguard/conn" + "github.com/sagernet/wireguard-go/conn" ) type Peer struct { diff --git a/device/queueconstants_android.go b/device/queueconstants_android.go index 236dea1..a3bee69 100644 --- a/device/queueconstants_android.go +++ b/device/queueconstants_android.go @@ -5,7 +5,7 @@ package device -import "golang.zx2c4.com/wireguard/conn" +import "github.com/sagernet/wireguard-go/conn" /* Reduce memory consumption for Android */ diff --git a/device/queueconstants_default.go b/device/queueconstants_default.go index b061185..1d09285 100644 --- a/device/queueconstants_default.go +++ b/device/queueconstants_default.go @@ -7,7 +7,7 @@ package device -import "golang.zx2c4.com/wireguard/conn" +import "github.com/sagernet/wireguard-go/conn" const ( QueueStagedSize = conn.IdealBatchSize diff --git a/device/receive.go b/device/receive.go index 1392957..2c40730 100644 --- a/device/receive.go +++ b/device/receive.go @@ -12,10 +12,10 @@ import ( "sync" "time" + "github.com/sagernet/wireguard-go/conn" "golang.org/x/crypto/chacha20poly1305" "golang.org/x/net/ipv4" "golang.org/x/net/ipv6" - "golang.zx2c4.com/wireguard/conn" ) type QueueHandshakeElement struct { @@ -411,7 +411,6 @@ func (device *Device) RoutineHandshake(id int) { // derive keypair err = peer.BeginSymmetricSession() - if err != nil { device.log.Errorf("%v - Failed to derive keypair: %v", peer, err) goto skip diff --git a/device/send.go b/device/send.go index ff8f7da..457fde7 100644 --- a/device/send.go +++ b/device/send.go @@ -13,11 +13,11 @@ import ( "sync" "time" + "github.com/sagernet/wireguard-go/conn" + "github.com/sagernet/wireguard-go/tun" "golang.org/x/crypto/chacha20poly1305" "golang.org/x/net/ipv4" "golang.org/x/net/ipv6" - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/tun" ) /* Outbound flow diff --git a/device/sticky_default.go b/device/sticky_default.go index 22e1e15..cac7add 100644 --- a/device/sticky_default.go +++ b/device/sticky_default.go @@ -3,8 +3,8 @@ package device import ( - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/rwcancel" + "github.com/sagernet/wireguard-go/conn" + "github.com/sagernet/wireguard-go/rwcancel" ) func (device *Device) startRouteListener(_ conn.Bind) (*rwcancel.RWCancel, error) { diff --git a/device/sticky_linux.go b/device/sticky_linux.go index f23ff02..9fcfeeb 100644 --- a/device/sticky_linux.go +++ b/device/sticky_linux.go @@ -18,10 +18,9 @@ import ( "sync" "unsafe" + "github.com/sagernet/wireguard-go/conn" + "github.com/sagernet/wireguard-go/rwcancel" "golang.org/x/sys/unix" - - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/rwcancel" ) func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, error) { diff --git a/device/tun.go b/device/tun.go index c85dd50..01a92ed 100644 --- a/device/tun.go +++ b/device/tun.go @@ -8,7 +8,7 @@ package device import ( "fmt" - "golang.zx2c4.com/wireguard/tun" + "github.com/sagernet/wireguard-go/tun" ) const DefaultMTU = 1420 diff --git a/device/uapi.go b/device/uapi.go index cc69488..cba371d 100644 --- a/device/uapi.go +++ b/device/uapi.go @@ -18,7 +18,7 @@ import ( "sync" "time" - "golang.zx2c4.com/wireguard/ipc" + "github.com/sagernet/wireguard-go/ipc" ) type IPCError struct { diff --git a/go.mod b/go.mod index 85766cf..eab479d 100644 --- a/go.mod +++ b/go.mod @@ -1,4 +1,4 @@ -module golang.zx2c4.com/wireguard +module github.com/sagernet/wireguard-go go 1.23.1 diff --git a/ipc/uapi_linux.go b/ipc/uapi_linux.go index fddded0..be59e58 100644 --- a/ipc/uapi_linux.go +++ b/ipc/uapi_linux.go @@ -9,8 +9,8 @@ import ( "net" "os" + "github.com/sagernet/wireguard-go/rwcancel" "golang.org/x/sys/unix" - "golang.zx2c4.com/wireguard/rwcancel" ) type UAPIListener struct { @@ -84,7 +84,6 @@ func UAPIListen(name string, file *os.File) (net.Listener, error) { unix.IN_DELETE| unix.IN_DELETE_SELF, ) - if err != nil { return nil, err } diff --git a/ipc/uapi_windows.go b/ipc/uapi_windows.go index 86e60b0..a146f1a 100644 --- a/ipc/uapi_windows.go +++ b/ipc/uapi_windows.go @@ -8,8 +8,8 @@ package ipc import ( "net" + "github.com/sagernet/wireguard-go/ipc/namedpipe" "golang.org/x/sys/windows" - "golang.zx2c4.com/wireguard/ipc/namedpipe" ) // TODO: replace these with actual standard windows error numbers from the win package diff --git a/tun/errors.go b/tun/errors.go index 75ae3a4..2c49fc7 100644 --- a/tun/errors.go +++ b/tun/errors.go @@ -4,9 +4,7 @@ import ( "errors" ) -var ( - // ErrTooManySegments is returned by Device.Read() when segmentation - // overflows the length of supplied buffers. This error should not cause - // reads to cease. - ErrTooManySegments = errors.New("too many segments") -) +// ErrTooManySegments is returned by Device.Read() when segmentation +// overflows the length of supplied buffers. This error should not cause +// reads to cease. +var ErrTooManySegments = errors.New("too many segments") diff --git a/tun/offload_linux.go b/tun/offload_linux.go index 5f0db06..d360d40 100644 --- a/tun/offload_linux.go +++ b/tun/offload_linux.go @@ -12,8 +12,8 @@ import ( "io" "unsafe" + "github.com/sagernet/wireguard-go/conn" "golang.org/x/sys/unix" - "golang.zx2c4.com/wireguard/conn" ) const tcpFlagsOffset = 13 diff --git a/tun/tun_linux.go b/tun/tun_linux.go index 1461e06..b10ec3c 100644 --- a/tun/tun_linux.go +++ b/tun/tun_linux.go @@ -17,9 +17,9 @@ import ( "time" "unsafe" + "github.com/sagernet/wireguard-go/conn" + "github.com/sagernet/wireguard-go/rwcancel" "golang.org/x/sys/unix" - "golang.zx2c4.com/wireguard/conn" - "golang.zx2c4.com/wireguard/rwcancel" ) const ( @@ -514,9 +514,7 @@ func (tun *NativeTun) initFromFlags(name string) error { return err } if e := sc.Control(func(fd uintptr) { - var ( - ifr *unix.Ifreq - ) + var ifr *unix.Ifreq ifr, err = unix.NewIfreq(name) if err != nil { return From 45cd03b8a17cbe11e20af9d55753a8edacf72b82 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 17 May 2026 19:34:11 +0800 Subject: [PATCH 24/34] Downgrade dependencies --- go.mod | 8 ++++---- go.sum | 12 ++++++------ 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/go.mod b/go.mod index eab479d..2dfe133 100644 --- a/go.mod +++ b/go.mod @@ -1,10 +1,10 @@ module github.com/sagernet/wireguard-go -go 1.23.1 +go 1.20 require ( - golang.org/x/crypto v0.37.0 - golang.org/x/net v0.39.0 - golang.org/x/sys v0.32.0 + golang.org/x/crypto v0.13.0 + golang.org/x/net v0.15.0 + golang.org/x/sys v0.12.0 golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 ) diff --git a/go.sum b/go.sum index c9da2b3..ec4169f 100644 --- a/go.sum +++ b/go.sum @@ -1,8 +1,8 @@ -golang.org/x/crypto v0.37.0 h1:kJNSjF/Xp7kU0iB2Z+9viTPMW4EqqsrywMXLJOOsXSE= -golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc= -golang.org/x/net v0.39.0 h1:ZCu7HMWDxpXpaiKdhzIfaltL9Lp31x/3fCP11bc6/fY= -golang.org/x/net v0.39.0/go.mod h1:X7NRbYVEA+ewNkCNyJ513WmMdQ3BineSwVtN2zD/d+E= -golang.org/x/sys v0.32.0 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20= -golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= +golang.org/x/crypto v0.13.0 h1:mvySKfSWJ+UKUii46M40LOvyWfN0s2U+46/jDd0e6Ck= +golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc= +golang.org/x/net v0.15.0 h1:ugBLEUaxABaB5AJqW9enI0ACdci2RUd4eP51NTBvuJ8= +golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk= +golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o= +golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI= From 749db0015c62b64eae1c2fd449177d28a6ae28ef Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 17 May 2026 20:11:01 +0800 Subject: [PATCH 25/34] Apply Tailscale bind send headroom --- conn/bind_std.go | 9 +++++---- conn/bind_windows.go | 3 ++- conn/conn.go | 8 +++++--- device/constants.go | 6 +++--- device/noise-protocol.go | 15 ++++++++------- device/peer.go | 5 ++++- device/send.go | 36 +++++++++++++++++++++++------------- 7 files changed, 50 insertions(+), 32 deletions(-) diff --git a/conn/bind_std.go b/conn/bind_std.go index 80e8210..7cbf9a8 100644 --- a/conn/bind_std.go +++ b/conn/bind_std.go @@ -334,7 +334,7 @@ func (e ErrUDPGSODisabled) Unwrap() error { return e.RetryErr } -func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint) error { +func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint, offset int) error { s.mu.Lock() blackhole := s.blackhole4 conn := s.ipv4 @@ -377,7 +377,7 @@ func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint) error { ) retry: if offload { - n := coalesceMessages(ua, endpoint.(*StdNetEndpoint), bufs, *msgs, setGSOSize) + n := coalesceMessages(ua, endpoint.(*StdNetEndpoint), bufs, offset, *msgs, setGSOSize) err = s.send(conn, br, (*msgs)[:n]) if err != nil && offload && errShouldDisableUDPGSO(err) { offload = false @@ -394,7 +394,7 @@ retry: } else { for i := range bufs { (*msgs)[i].Addr = ua - (*msgs)[i].Buffers[0] = bufs[i] + (*msgs)[i].Buffers[0] = bufs[i][offset:] setSrcControl(&(*msgs)[i].OOB, endpoint.(*StdNetEndpoint)) } err = s.send(conn, br, (*msgs)[:len(bufs)]) @@ -443,7 +443,7 @@ const ( type setGSOFunc func(control *[]byte, gsoSize uint16) -func coalesceMessages(addr *net.UDPAddr, ep *StdNetEndpoint, bufs [][]byte, msgs []ipv6.Message, setGSO setGSOFunc) int { +func coalesceMessages(addr *net.UDPAddr, ep *StdNetEndpoint, bufs [][]byte, offset int, msgs []ipv6.Message, setGSO setGSOFunc) int { var ( base = -1 // index of msg we are currently coalescing into gsoSize int // segmentation size of msgs[base] @@ -455,6 +455,7 @@ func coalesceMessages(addr *net.UDPAddr, ep *StdNetEndpoint, bufs [][]byte, msgs maxPayloadLen = maxIPv6PayloadLen } for i, buf := range bufs { + buf = buf[offset:] if i > 0 { msgLen := len(buf) baseLenBefore := len(msgs[base].Buffers[0]) diff --git a/conn/bind_windows.go b/conn/bind_windows.go index d166227..66d91cd 100644 --- a/conn/bind_windows.go +++ b/conn/bind_windows.go @@ -485,7 +485,7 @@ func (bind *afWinRingBind) Send(buf []byte, nend *WinRingEndpoint, isOpen *atomi return winrio.SendEx(bind.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0) } -func (bind *WinRingBind) Send(bufs [][]byte, endpoint Endpoint) error { +func (bind *WinRingBind) Send(bufs [][]byte, endpoint Endpoint, offset int) error { nend, ok := endpoint.(*WinRingEndpoint) if !ok { return ErrWrongEndpointType @@ -493,6 +493,7 @@ func (bind *WinRingBind) Send(bufs [][]byte, endpoint Endpoint) error { bind.mu.RLock() defer bind.mu.RUnlock() for _, buf := range bufs { + buf = buf[offset:] switch nend.family { case windows.AF_INET: if bind.v4.blackhole { diff --git a/conn/conn.go b/conn/conn.go index 1304657..139b06e 100644 --- a/conn/conn.go +++ b/conn/conn.go @@ -45,9 +45,11 @@ type Bind interface { // This mark is passed to the kernel as the socket option SO_MARK. SetMark(mark uint32) error - // Send writes one or more packets in bufs to address ep. The length of - // bufs must not exceed BatchSize(). - Send(bufs [][]byte, ep Endpoint) error + // Send writes one or more packets in bufs to address ep. A nonzero offset + // can be used to instruct the Bind on where packet data begins in each + // element of the bufs slice. Space preceding offset is free to use for + // additional encapsulation. The length of bufs must not exceed BatchSize(). + Send(bufs [][]byte, ep Endpoint, offset int) error // ParseEndpoint creates a new endpoint from a string. ParseEndpoint(s string) (Endpoint, error) diff --git a/device/constants.go b/device/constants.go index 41da618..1a02daf 100644 --- a/device/constants.go +++ b/device/constants.go @@ -27,9 +27,9 @@ const ( ) const ( - MinMessageSize = MessageKeepaliveSize // minimum size of transport message (keepalive) - MaxMessageSize = MaxSegmentSize // maximum size of transport message - MaxContentSize = MaxSegmentSize - MessageTransportSize // maximum size of transport message content + MinMessageSize = MessageKeepaliveSize // minimum size of transport message (keepalive) + MaxMessageSize = MaxSegmentSize // maximum size of transport message + MaxContentSize = MaxSegmentSize - MessageTransportSize - MessageEncapsulatingTransportSize // maximum size of transport message content ) /* Implementation constants */ diff --git a/device/noise-protocol.go b/device/noise-protocol.go index a9b8498..98a03bc 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -60,13 +60,14 @@ 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 // 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 + MessageEncapsulatingTransportSize = 8 // size of optional, free (for use by conn.Bind.Send()) space preceding the transport header + MessageTransportSize = MessageTransportHeaderSize + poly1305.TagSize // size of empty transport + MessageKeepaliveSize = MessageTransportSize // size of keepalive + MessageHandshakeSize = MessageInitiationSize // size of largest handshake related message ) const ( diff --git a/device/peer.go b/device/peer.go index fff1cf5..bca121d 100644 --- a/device/peer.go +++ b/device/peer.go @@ -113,6 +113,9 @@ func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) { return peer, nil } +// SendBuffers sends buffers to peer. WireGuard packet data in each element of +// buffers must be preceded by MessageEncapsulatingTransportSize number of +// bytes. func (peer *Peer) SendBuffers(buffers [][]byte) error { peer.device.net.RLock() defer peer.device.net.RUnlock() @@ -133,7 +136,7 @@ func (peer *Peer) SendBuffers(buffers [][]byte) error { } peer.endpoint.Unlock() - err := peer.device.net.bind.Send(buffers, endpoint) + err := peer.device.net.bind.Send(buffers, endpoint, MessageEncapsulatingTransportSize) if err == nil { var totalLen uint64 for _, b := range buffers { diff --git a/device/send.go b/device/send.go index 457fde7..78a3108 100644 --- a/device/send.go +++ b/device/send.go @@ -45,11 +45,15 @@ import ( */ type QueueOutboundElement struct { - buffer *[MaxMessageSize]byte // slice holding the packet data - packet []byte // slice of "buffer" (always!) - nonce uint64 // nonce for encryption - keypair *Keypair // keypair for encryption - peer *Peer // related peer + buffer *[MaxMessageSize]byte // slice holding the packet data + // packet is always a slice of "buffer". The starting offset in buffer + // is either: + // a) MessageEncapsulatingTransportSize+MessageTransportHeaderSize (plaintext) + // b) 0 (post-encryption) + packet []byte + nonce uint64 // nonce for encryption + keypair *Keypair // keypair for encryption + peer *Peer // related peer } type QueueOutboundElementsContainer struct { @@ -123,14 +127,15 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error { return err } - packet := make([]byte, MessageInitiationSize) + buf := make([]byte, MessageEncapsulatingTransportSize+MessageInitiationSize) + packet := buf[MessageEncapsulatingTransportSize:] _ = msg.marshal(packet) peer.cookieGenerator.AddMacs(packet) peer.timersAnyAuthenticatedPacketTraversal() peer.timersAnyAuthenticatedPacketSent() - err = peer.SendBuffers([][]byte{packet}) + err = peer.SendBuffers([][]byte{buf}) if err != nil { peer.device.log.Errorf("%v - Failed to send handshake initiation: %v", peer, err) } @@ -152,7 +157,8 @@ func (peer *Peer) SendHandshakeResponse() error { return err } - packet := make([]byte, MessageResponseSize) + buf := make([]byte, MessageEncapsulatingTransportSize+MessageResponseSize) + packet := buf[MessageEncapsulatingTransportSize:] _ = response.marshal(packet) peer.cookieGenerator.AddMacs(packet) @@ -167,7 +173,7 @@ func (peer *Peer) SendHandshakeResponse() error { peer.timersAnyAuthenticatedPacketSent() // TODO: allocation could be avoided - err = peer.SendBuffers([][]byte{packet}) + err = peer.SendBuffers([][]byte{buf}) if err != nil { peer.device.log.Errorf("%v - Failed to send handshake response: %v", peer, err) } @@ -184,10 +190,11 @@ func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement) return err } - packet := make([]byte, MessageCookieReplySize) + buf := make([]byte, MessageEncapsulatingTransportSize+MessageCookieReplySize) + packet := buf[MessageEncapsulatingTransportSize:] _ = reply.marshal(packet) // TODO: allocation could be avoided - device.net.bind.Send([][]byte{packet}, initiatingElem.endpoint) + device.net.bind.Send([][]byte{buf}, initiatingElem.endpoint, MessageEncapsulatingTransportSize) return nil } @@ -220,7 +227,7 @@ func (device *Device) RoutineReadFromTUN() { elemsByPeer = make(map[*Peer]*QueueOutboundElementsContainer, batchSize) count = 0 sizes = make([]int, batchSize) - offset = MessageTransportHeaderSize + offset = MessageEncapsulatingTransportSize + MessageTransportHeaderSize ) for i := range elems { @@ -446,7 +453,7 @@ func (device *Device) RoutineEncryption(id int) { for elemsContainer := range device.queue.encryption.c { for _, elem := range elemsContainer.elems { // populate header fields - header := elem.buffer[:MessageTransportHeaderSize] + header := elem.buffer[MessageEncapsulatingTransportSize : MessageEncapsulatingTransportSize+MessageTransportHeaderSize] fieldType := header[0:4] fieldReceiver := header[4:8] @@ -469,6 +476,9 @@ func (device *Device) RoutineEncryption(id int) { elem.packet, nil, ) + + // re-slice packet to include encapsulating transport space + elem.packet = elem.buffer[:MessageEncapsulatingTransportSize+len(elem.packet)] } elemsContainer.Unlock() } From 414291f6d6ab8a4041e78b400be0f3af0c0b1d7b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 17 May 2026 20:12:11 +0800 Subject: [PATCH 26/34] Apply Tailscale endpoint awareness --- conn/conn.go | 34 ++++++++++++++++++++++++++++++++++ device/noise-protocol.go | 8 +++++++- device/receive.go | 5 ++++- 3 files changed, 45 insertions(+), 2 deletions(-) diff --git a/conn/conn.go b/conn/conn.go index 139b06e..87b7510 100644 --- a/conn/conn.go +++ b/conn/conn.go @@ -86,6 +86,40 @@ type Endpoint interface { SrcIP() netip.Addr } +// InitiationAwareEndpoint is an optional [Endpoint] specialization for +// integrations that want to know when a WireGuard handshake initiation +// message has been received, enabling just-in-time peer configuration before +// attempted decryption. +// +// It's most useful when used in combination with [PeerAwareEndpoint], enabling +// JIT peer configuration and post-decryption peer verification from a single +// implementer. +type InitiationAwareEndpoint interface { + // InitiationMessagePublicKey is called when a handshake initiation message + // has been received, and the sender's public key has been identified, but + // BEFORE an attempt has been made to verify it. + InitiationMessagePublicKey(peerPublicKey [32]byte) +} + +// PeerAwareEndpoint is an optional Endpoint specialization for +// integrations that want to know about the outcome of Cryptokey Routing +// identification. +// +// If they receive a packet from a source they had not pre-identified, +// to learn the identification WireGuard can derive from the session +// or handshake. +// +// A [PeerAwareEndpoint] may be installed as the [conn.Endpoint] following +// successful decryption unless endpoint roaming has been disabled for +// the peer. +type PeerAwareEndpoint interface { + // FromPeer is called at least once per successfully Cryptokey Routing ID'd + // [ReceiveFunc] packets batch for a given node key. wireguard-go will + // always call it for the latest/tail packet in the batch, only ever + // suppressing calls for older packets. + FromPeer(peerPublicKey [32]byte) +} + var ( ErrBindAlreadyOpen = errors.New("bind is already open") ErrWrongEndpointType = errors.New("endpoint type does not correspond with bind type") diff --git a/device/noise-protocol.go b/device/noise-protocol.go index 98a03bc..ed4d82a 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -12,6 +12,7 @@ import ( "sync" "time" + "github.com/sagernet/wireguard-go/conn" "github.com/sagernet/wireguard-go/tai64n" "golang.org/x/crypto/blake2s" "golang.org/x/crypto/chacha20poly1305" @@ -337,7 +338,7 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e return &msg, nil } -func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { +func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation, endpoint conn.Endpoint) *Peer { var ( hash [blake2s.Size]byte chainKey [blake2s.Size]byte @@ -371,6 +372,11 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer { // lookup peer + initEP, ok := endpoint.(conn.InitiationAwareEndpoint) + if ok { + initEP.InitiationMessagePublicKey(peerPK) + } + peer := device.LookupPeer(peerPK) if peer == nil || !peer.isRunning.Load() { return nil diff --git a/device/receive.go b/device/receive.go index 2c40730..e13c987 100644 --- a/device/receive.go +++ b/device/receive.go @@ -359,7 +359,7 @@ func (device *Device) RoutineHandshake(id int) { // consume initiation - peer := device.ConsumeMessageInitiation(&msg) + peer := device.ConsumeMessageInitiation(&msg, elem.endpoint) if peer == nil { device.log.Verbosef("Received invalid initiation message from %s", elem.endpoint.DstToString()) goto skip @@ -459,6 +459,9 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) { peer.timersHandshakeComplete() peer.SendStagedPackets() } + if ep, ok := elem.endpoint.(conn.PeerAwareEndpoint); ok { + ep.FromPeer(peer.handshake.remoteStatic) + } rxBytesLen += uint64(len(elem.packet) + MinMessageSize) if len(elem.packet) == 0 { From 824e7573c086a09674a7a6c90061664719df8395 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 17 May 2026 20:15:11 +0800 Subject: [PATCH 27/34] Apply Tailscale TUN offload APIs --- tun/checksum.go | 11 +++ tun/offload.go | 200 +++++++++++++++++++++++++++++++++++++++++++ tun/offload_linux.go | 18 ++-- tun/tun.go | 14 +++ tun/tun_linux.go | 47 +++++++++- 5 files changed, 277 insertions(+), 13 deletions(-) create mode 100644 tun/offload.go diff --git a/tun/checksum.go b/tun/checksum.go index b489c56..ac16569 100644 --- a/tun/checksum.go +++ b/tun/checksum.go @@ -92,6 +92,11 @@ func checksum(b []byte, initial uint64) uint16 { return uint16(ac) } +// Checksum computes an IP checksum starting with the provided initial value. +func Checksum(data []byte, initial uint16) uint16 { + return checksum(data, uint64(initial)) +} + func pseudoHeaderChecksumNoFold(protocol uint8, srcAddr, dstAddr []byte, totalLen uint16) uint64 { sum := checksumNoFold(srcAddr, 0) sum = checksumNoFold(dstAddr, sum) @@ -100,3 +105,9 @@ func pseudoHeaderChecksumNoFold(protocol uint8, srcAddr, dstAddr []byte, totalLe binary.BigEndian.PutUint16(tmp, totalLen) return checksumNoFold(tmp, sum) } + +// PseudoHeaderChecksum computes an IP pseudo-header checksum. srcAddr and +// dstAddr must be 4 or 16 bytes in length. +func PseudoHeaderChecksum(protocol uint8, srcAddr, dstAddr []byte, totalLen uint16) uint16 { + return checksum([]byte{}, pseudoHeaderChecksumNoFold(protocol, srcAddr, dstAddr, totalLen)) +} diff --git a/tun/offload.go b/tun/offload.go new file mode 100644 index 0000000..4e84db4 --- /dev/null +++ b/tun/offload.go @@ -0,0 +1,200 @@ +package tun + +import ( + "encoding/binary" + "fmt" +) + +// GSOType represents the type of segmentation offload. +type GSOType int + +const ( + GSONone GSOType = iota + GSOTCPv4 + GSOTCPv6 + GSOUDPL4 +) + +func (g GSOType) String() string { + switch g { + case GSONone: + return "GSONone" + case GSOTCPv4: + return "GSOTCPv4" + case GSOTCPv6: + return "GSOTCPv6" + case GSOUDPL4: + return "GSOUDPL4" + default: + return "unknown" + } +} + +// GSOOptions is loosely modeled after struct virtio_net_hdr from the VIRTIO +// specification. It is a common representation of GSO metadata that can be +// applied to support packet GSO across tun.Device implementations. +type GSOOptions struct { + // GSOType represents the type of segmentation offload. + GSOType GSOType + // HdrLen is the sum of the layer 3 and 4 header lengths. This field may be + // zero when GSOType == GSONone. + HdrLen uint16 + // CsumStart is the head byte index of the packet data to be checksummed, + // i.e. the start of the TCP or UDP header. + CsumStart uint16 + // CsumOffset is the offset from CsumStart where the 2-byte checksum value + // should be placed. + CsumOffset uint16 + // GSOSize is the size of each segment exclusive of HdrLen. The tail segment + // may be smaller than this value. + GSOSize uint16 + // NeedsCsum may be set where GSOType == GSONone. When set, the checksum + // at CsumStart + CsumOffset must be a partial checksum, i.e. the + // pseudo-header sum. + NeedsCsum bool +} + +const ( + gsoIPv4SrcAddrOffset = 12 + gsoIPv6SrcAddrOffset = 8 + gsoTCPFlagsOffset = 13 + gsoIPProtoTCP = 6 + gsoIPProtoUDP = 17 +) + +const ( + gsoTCPFlagFIN uint8 = 0x01 + gsoTCPFlagPSH uint8 = 0x08 +) + +// GSOSplit splits packets from in into outBufs[][outOffset:], writing +// the size of each element into sizes. It returns the number of buffers +// populated, and/or an error. Callers may pass an in slice that overlaps with +// the first element of outBufs, i.e. &in[0] may be equal to +// &outBufs[0][outOffset]. GSONone is a valid options.GSOType regardless of the +// value of options.NeedsCsum. Length of each outBufs element must be greater +// than or equal to the length of in, otherwise output may be silently truncated. +func GSOSplit(in []byte, options GSOOptions, outBufs [][]byte, sizes []int, outOffset int) (int, error) { + cSumAt := int(options.CsumStart) + int(options.CsumOffset) + if cSumAt+1 >= len(in) { + return 0, fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(in)) + } + + if len(in) < int(options.HdrLen) { + return 0, fmt.Errorf("length of packet (%d) < GSO HdrLen (%d)", len(in), options.HdrLen) + } + + payloadLen := len(in) - int(options.HdrLen) + if options.GSOType == GSONone || payloadLen < int(options.GSOSize) { + if len(in) > len(outBufs[0][outOffset:]) { + return 0, fmt.Errorf("length of packet (%d) exceeds output element length (%d)", len(in), len(outBufs[0][outOffset:])) + } + if options.NeedsCsum { + initial := binary.BigEndian.Uint16(in[cSumAt:]) + in[cSumAt], in[cSumAt+1] = 0, 0 + binary.BigEndian.PutUint16(in[cSumAt:], ^Checksum(in[options.CsumStart:], initial)) + } + sizes[0] = copy(outBufs[0][outOffset:], in) + return 1, nil + } + + if options.HdrLen < options.CsumStart { + return 0, fmt.Errorf("GSO HdrLen (%d) < GSO CsumStart (%d)", options.HdrLen, options.CsumStart) + } + + ipVersion := in[0] >> 4 + switch ipVersion { + case 4: + if options.GSOType != GSOTCPv4 && options.GSOType != GSOUDPL4 { + return 0, fmt.Errorf("ip header version: %d, GSO type: %s", ipVersion, options.GSOType) + } + if len(in) < 20 { + return 0, fmt.Errorf("length of packet (%d) < minimum ipv4 header size (%d)", len(in), 20) + } + case 6: + if options.GSOType != GSOTCPv6 && options.GSOType != GSOUDPL4 { + return 0, fmt.Errorf("ip header version: %d, GSO type: %s", ipVersion, options.GSOType) + } + if len(in) < 40 { + return 0, fmt.Errorf("length of packet (%d) < minimum ipv6 header size (%d)", len(in), 40) + } + default: + return 0, fmt.Errorf("invalid ip header version: %d", ipVersion) + } + + iphLen := int(options.CsumStart) + srcAddrOffset := gsoIPv6SrcAddrOffset + addrLen := 16 + if ipVersion == 4 { + srcAddrOffset = gsoIPv4SrcAddrOffset + addrLen = 4 + } + transportCsumAt := int(options.CsumStart + options.CsumOffset) + var firstTCPSeqNum uint32 + var protocol uint8 + if options.GSOType == GSOTCPv4 || options.GSOType == GSOTCPv6 { + protocol = gsoIPProtoTCP + if len(in) < int(options.CsumStart)+20 { + return 0, fmt.Errorf("length of packet (%d) < GSO CsumStart (%d) + minimum TCP header size (%d)", + len(in), options.CsumStart, 20) + } + firstTCPSeqNum = binary.BigEndian.Uint32(in[options.CsumStart+4:]) + } else { + protocol = gsoIPProtoUDP + } + nextSegmentDataAt := int(options.HdrLen) + i := 0 + for ; nextSegmentDataAt < len(in); i++ { + if i == len(outBufs) { + return i - 1, ErrTooManySegments + } + nextSegmentEnd := nextSegmentDataAt + int(options.GSOSize) + if nextSegmentEnd > len(in) { + nextSegmentEnd = len(in) + } + segmentDataLen := nextSegmentEnd - nextSegmentDataAt + totalLen := int(options.HdrLen) + segmentDataLen + sizes[i] = totalLen + out := outBufs[i][outOffset:] + + copy(out, in[:iphLen]) + if ipVersion == 4 { + if i > 0 { + id := binary.BigEndian.Uint16(out[4:]) + id += uint16(i) + binary.BigEndian.PutUint16(out[4:], id) + } + out[10], out[11] = 0, 0 + binary.BigEndian.PutUint16(out[2:], uint16(totalLen)) + ipv4CSum := ^Checksum(out[:iphLen], 0) + binary.BigEndian.PutUint16(out[10:], ipv4CSum) + } else { + binary.BigEndian.PutUint16(out[4:], uint16(totalLen-iphLen)) + } + + copy(out[options.CsumStart:options.HdrLen], in[options.CsumStart:options.HdrLen]) + + if protocol == gsoIPProtoTCP { + tcpSeq := firstTCPSeqNum + uint32(options.GSOSize*uint16(i)) + binary.BigEndian.PutUint32(out[options.CsumStart+4:], tcpSeq) + if nextSegmentEnd != len(in) { + clearFlags := gsoTCPFlagFIN | gsoTCPFlagPSH + out[options.CsumStart+gsoTCPFlagsOffset] &^= clearFlags + } + } else { + binary.BigEndian.PutUint16(out[options.CsumStart+4:], uint16(segmentDataLen)+(options.HdrLen-options.CsumStart)) + } + + copy(out[options.HdrLen:], in[nextSegmentDataAt:nextSegmentEnd]) + + out[transportCsumAt], out[transportCsumAt+1] = 0, 0 + transportHeaderLen := int(options.HdrLen - options.CsumStart) + lenForPseudo := uint16(transportHeaderLen + segmentDataLen) + transportCSum := PseudoHeaderChecksum(protocol, in[srcAddrOffset:srcAddrOffset+addrLen], in[srcAddrOffset+addrLen:srcAddrOffset+addrLen*2], lenForPseudo) + transportCSum = ^Checksum(out[options.CsumStart:totalLen], transportCSum) + binary.BigEndian.PutUint16(out[options.CsumStart+options.CsumOffset:], transportCSum) + + nextSegmentDataAt += int(options.GSOSize) + } + return i, nil +} diff --git a/tun/offload_linux.go b/tun/offload_linux.go index d360d40..6825bbc 100644 --- a/tun/offload_linux.go +++ b/tun/offload_linux.go @@ -748,7 +748,7 @@ const ( udp6GROCandidate ) -func packetIsGROCandidate(b []byte, canUDPGRO bool) groCandidateType { +func packetIsGROCandidate(b []byte, gro groDisablementFlags) groCandidateType { if len(b) < 28 { return notGROCandidate } @@ -757,17 +757,17 @@ func packetIsGROCandidate(b []byte, canUDPGRO bool) groCandidateType { // IPv4 packets w/IP options do not coalesce return notGROCandidate } - if b[9] == unix.IPPROTO_TCP && len(b) >= 40 { + if b[9] == unix.IPPROTO_TCP && len(b) >= 40 && gro.canTCPGRO() { return tcp4GROCandidate } - if b[9] == unix.IPPROTO_UDP && canUDPGRO { + if b[9] == unix.IPPROTO_UDP && gro.canUDPGRO() { return udp4GROCandidate } } else if b[0]>>4 == 6 { - if b[6] == unix.IPPROTO_TCP && len(b) >= 60 { + if b[6] == unix.IPPROTO_TCP && len(b) >= 60 && gro.canTCPGRO() { return tcp6GROCandidate } - if b[6] == unix.IPPROTO_UDP && len(b) >= 48 && canUDPGRO { + if b[6] == unix.IPPROTO_UDP && len(b) >= 48 && gro.canUDPGRO() { return udp6GROCandidate } } @@ -860,15 +860,15 @@ func udpGRO(bufs [][]byte, offset int, pktI int, table *udpGROTable, isV6 bool) // handleGRO evaluates bufs for GRO, and writes the indices of the resulting // packets into toWrite. toWrite, tcpTable, and udpTable should initially be // empty (but non-nil), and are passed in to save allocs as the caller may reset -// and recycle them across vectors of packets. canUDPGRO indicates if UDP GRO is -// supported. -func handleGRO(bufs [][]byte, offset int, tcpTable *tcpGROTable, udpTable *udpGROTable, canUDPGRO bool, toWrite *[]int) error { +// and recycle them across vectors of packets. gro indicates if TCP and UDP GRO +// are supported/enabled. +func handleGRO(bufs [][]byte, offset int, tcpTable *tcpGROTable, udpTable *udpGROTable, gro groDisablementFlags, toWrite *[]int) error { for i := range bufs { if offset < virtioNetHdrLen || offset > len(bufs[i])-1 { return errors.New("invalid offset") } var result groResult - switch packetIsGROCandidate(bufs[i][offset:], canUDPGRO) { + switch packetIsGROCandidate(bufs[i][offset:], gro) { case tcp4GROCandidate: result = tcpGRO(bufs, offset, i, tcpTable, false) case tcp6GROCandidate: diff --git a/tun/tun.go b/tun/tun.go index 336d642..6fb5c56 100644 --- a/tun/tun.go +++ b/tun/tun.go @@ -51,3 +51,17 @@ type Device interface { // lifetime of a Device. BatchSize() int } + +// GRODevice is a Device extended with methods for disabling GRO. Certain OS +// versions may have offload bugs. Where these bugs negatively impact throughput +// or break connectivity entirely we can use these methods to disable the +// related offload. +type GRODevice interface { + Device + + // DisableUDPGRO disables UDP GRO if it is enabled. + DisableUDPGRO() + + // DisableTCPGRO disables TCP GRO if it is enabled. + DisableTCPGRO() +} diff --git a/tun/tun_linux.go b/tun/tun_linux.go index b10ec3c..4b7866d 100644 --- a/tun/tun_linux.go +++ b/tun/tun_linux.go @@ -38,7 +38,6 @@ type NativeTun struct { statusListenersShutdown chan struct{} batchSize int vnetHdr bool - udpGSO bool closeOnce sync.Once @@ -49,10 +48,34 @@ type NativeTun struct { readOpMu sync.Mutex // readOpMu guards readBuff readBuff [virtioNetHdrLen + 65535]byte // if vnetHdr every read() is prefixed by virtioNetHdr - writeOpMu sync.Mutex // writeOpMu guards toWrite, tcpGROTable + writeOpMu sync.Mutex // writeOpMu guards toWrite, tcpGROTable, udpGROTable, gro toWrite []int tcpGROTable *tcpGROTable udpGROTable *udpGROTable + gro groDisablementFlags +} + +type groDisablementFlags int + +const ( + tcpGRODisabled groDisablementFlags = 1 << iota + udpGRODisabled +) + +func (g *groDisablementFlags) disableTCPGRO() { + *g |= tcpGRODisabled +} + +func (g *groDisablementFlags) canTCPGRO() bool { + return (*g)&tcpGRODisabled == 0 +} + +func (g *groDisablementFlags) disableUDPGRO() { + *g |= udpGRODisabled +} + +func (g *groDisablementFlags) canUDPGRO() bool { + return (*g)&udpGRODisabled == 0 } func (tun *NativeTun) File() *os.File { @@ -345,7 +368,7 @@ func (tun *NativeTun) Write(bufs [][]byte, offset int) (int, error) { ) tun.toWrite = tun.toWrite[:0] if tun.vnetHdr { - err := handleGRO(bufs, offset, tun.tcpGROTable, tun.udpGROTable, tun.udpGSO, &tun.toWrite) + err := handleGRO(bufs, offset, tun.tcpGROTable, tun.udpGROTable, tun.gro, &tun.toWrite) if err != nil { return 0, err } @@ -502,6 +525,20 @@ func (tun *NativeTun) BatchSize() int { return tun.batchSize } +// DisableUDPGRO disables UDP GRO if it is enabled. +func (tun *NativeTun) DisableUDPGRO() { + tun.writeOpMu.Lock() + tun.gro.disableUDPGRO() + tun.writeOpMu.Unlock() +} + +// DisableTCPGRO disables TCP GRO if it is enabled. +func (tun *NativeTun) DisableTCPGRO() { + tun.writeOpMu.Lock() + tun.gro.disableTCPGRO() + tun.writeOpMu.Unlock() +} + const ( // TODO: support TSO with ECN bits tunTCPOffloads = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 @@ -535,7 +572,9 @@ func (tun *NativeTun) initFromFlags(name string) error { tun.batchSize = conn.IdealBatchSize // tunUDPOffloads were added in Linux v6.2. We do not return an // error if they are unsupported at runtime. - tun.udpGSO = unix.IoctlSetInt(int(fd), unix.TUNSETOFFLOAD, tunTCPOffloads|tunUDPOffloads) == nil + if unix.IoctlSetInt(int(fd), unix.TUNSETOFFLOAD, tunTCPOffloads|tunUDPOffloads) != nil { + tun.gro.disableUDPGRO() + } } else { tun.batchSize = 1 } From b4db0692d3ad34898a95be7ec4ef59b14861c336 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Mon, 7 Nov 2022 15:41:48 +0800 Subject: [PATCH 28/34] Add custom worker size params (cherry picked from commit 7c2acadba17cadf8a1df957c49e1333130d460ad) (cherry picked from commit a7bac1754e7717e1d4009d1ffd2d13330067d631) (cherry picked from commit 7a2f11c693b49e784318bbf987173095c67b563d) --- device/device.go | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/device/device.go b/device/device.go index bfc56bb..88cdf61 100644 --- a/device/device.go +++ b/device/device.go @@ -281,7 +281,7 @@ func (device *Device) SetPrivateKey(sk NoisePrivateKey) error { return nil } -func NewDevice(tunDevice tun.Device, bind conn.Bind, logger *Logger) *Device { +func NewDevice(tunDevice tun.Device, bind conn.Bind, logger *Logger, workers int) *Device { device := new(Device) device.state.state.Store(uint32(deviceStateDown)) device.closed = make(chan struct{}) @@ -308,10 +308,12 @@ func NewDevice(tunDevice tun.Device, bind conn.Bind, logger *Logger) *Device { // start workers - cpus := runtime.NumCPU() + if workers == 0 { + workers = runtime.NumCPU() + } device.state.stopping.Wait() - device.queue.encryption.wg.Add(cpus) // One for each RoutineHandshake - for i := 0; i < cpus; i++ { + device.queue.encryption.wg.Add(workers) // One for each RoutineHandshake + for i := 0; i < workers; i++ { go device.RoutineEncryption(i + 1) go device.RoutineDecryption(i + 1) go device.RoutineHandshake(i + 1) From 7be452de15404412946d6a0d1706bd5cced956e8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Mon, 7 Aug 2023 20:57:31 +0800 Subject: [PATCH 29/34] Add pause support --- device/device.go | 13 +++++++++---- device/timers.go | 3 +++ go.mod | 3 ++- go.sum | 6 ++++-- 4 files changed, 18 insertions(+), 7 deletions(-) diff --git a/device/device.go b/device/device.go index 88cdf61..8afa6fd 100644 --- a/device/device.go +++ b/device/device.go @@ -6,11 +6,14 @@ package device import ( + "context" "runtime" "sync" "sync/atomic" "time" + "github.com/sagernet/sing/service" + "github.com/sagernet/sing/service/pause" "github.com/sagernet/wireguard-go/conn" "github.com/sagernet/wireguard-go/ratelimiter" "github.com/sagernet/wireguard-go/rwcancel" @@ -86,9 +89,10 @@ type Device struct { mtu atomic.Int32 } - ipcMutex sync.RWMutex - closed chan struct{} - log *Logger + ipcMutex sync.RWMutex + closed chan struct{} + log *Logger + pauseManager pause.Manager } // deviceState represents the state of a Device. @@ -281,8 +285,9 @@ func (device *Device) SetPrivateKey(sk NoisePrivateKey) error { return nil } -func NewDevice(tunDevice tun.Device, bind conn.Bind, logger *Logger, workers int) *Device { +func NewDevice(ctx context.Context, tunDevice tun.Device, bind conn.Bind, logger *Logger, workers int) *Device { device := new(Device) + device.pauseManager = service.FromContext[pause.Manager](ctx) device.state.state.Store(uint32(deviceStateDown)) device.closed = make(chan struct{}) device.log = logger diff --git a/device/timers.go b/device/timers.go index 32519aa..80fb7d9 100644 --- a/device/timers.go +++ b/device/timers.go @@ -39,6 +39,9 @@ func (peer *Peer) NewTimer(expirationFunction func(*Peer)) *Timer { timer.isPending = false timer.modifyingLock.Unlock() + if pauseManager := peer.device.pauseManager; pauseManager != nil { + pauseManager.WaitActive() + } expirationFunction(peer) }) timer.Stop() diff --git a/go.mod b/go.mod index 2dfe133..773bf00 100644 --- a/go.mod +++ b/go.mod @@ -3,8 +3,9 @@ module github.com/sagernet/wireguard-go go 1.20 require ( + github.com/sagernet/sing v0.7.10 golang.org/x/crypto v0.13.0 golang.org/x/net v0.15.0 - golang.org/x/sys v0.12.0 + golang.org/x/sys v0.21.0 golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 ) diff --git a/go.sum b/go.sum index ec4169f..9ce9725 100644 --- a/go.sum +++ b/go.sum @@ -1,8 +1,10 @@ +github.com/sagernet/sing v0.7.10 h1:2yPhZFx+EkyHPH8hXNezgyRSHyGY12CboId7CtwLROw= +github.com/sagernet/sing v0.7.10/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak= golang.org/x/crypto v0.13.0 h1:mvySKfSWJ+UKUii46M40LOvyWfN0s2U+46/jDd0e6Ck= golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc= golang.org/x/net v0.15.0 h1:ugBLEUaxABaB5AJqW9enI0ACdci2RUd4eP51NTBvuJ8= golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk= -golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o= -golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.21.0 h1:rF+pYz3DAGSQAxAu1CbC7catZg4ebC4UIeIhKxBZvws= +golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI= From a71256d250b1838f2db96db7a80a68a1a06d40de Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Wed, 20 Nov 2024 18:53:51 +0800 Subject: [PATCH 30/34] Add device.InputPacket --- device/send.go | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/device/send.go b/device/send.go index 78a3108..64a4221 100644 --- a/device/send.go +++ b/device/send.go @@ -322,6 +322,30 @@ func (device *Device) RoutineReadFromTUN() { } } +func (device *Device) InputPacket(destination []byte, packetSlices [][]byte) { + peer := device.allowedips.Lookup(destination) + if peer == nil { + return + } + elem := device.NewOutboundElement() + packet := elem.buffer[MessageTransportHeaderSize:] + var n int + for _, packetSlice := range packetSlices { + n += copy(packet[n:], packetSlice) + } + elem.packet = packet[:n] + elemsForPeer := device.GetOutboundElementsContainer() + if peer.isRunning.Load() { + elemsForPeer.elems = append(elemsForPeer.elems, elem) + peer.StagePackets(elemsForPeer) + peer.SendStagedPackets() + } else { + device.PutMessageBuffer(elem.buffer) + device.PutOutboundElement(elem) + device.PutOutboundElementsContainer(elemsForPeer) + } +} + func (peer *Peer) StagePackets(elems *QueueOutboundElementsContainer) { for { select { From 73f8c6542b658e44eac4b98b539c5741ed208599 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Fri, 15 Dec 2023 18:44:05 +0800 Subject: [PATCH 31/34] Export std net bind --- conn/bind_std.go | 58 +++++++++++++++++++++++++++++---- conn/bind_windows.go | 76 ++++++++++++++++++++++++++++++++++++++++---- conn/conn.go | 2 ++ conn/controlfns.go | 26 ++------------- conn/default.go | 6 +++- 5 files changed, 129 insertions(+), 39 deletions(-) diff --git a/conn/bind_std.go b/conn/bind_std.go index 7cbf9a8..cb979e3 100644 --- a/conn/bind_std.go +++ b/conn/bind_std.go @@ -16,6 +16,9 @@ import ( "sync" "syscall" + "github.com/sagernet/sing/common" + "github.com/sagernet/sing/common/control" + M "github.com/sagernet/sing/common/metadata" "golang.org/x/net/ipv4" "golang.org/x/net/ipv6" ) @@ -28,6 +31,9 @@ var _ Bind = (*StdNetBind)(nil) // methods for sending and receiving multiple datagrams per-syscall. See the // proposal in https://github.com/golang/go/issues/45886#issuecomment-1218301564. type StdNetBind struct { + externalControl control.Func + reservedForEndpoint map[netip.AddrPort][3]uint8 + mu sync.Mutex // protects all fields except as specified ipv4 *net.UDPConn ipv6 *net.UDPConn @@ -46,8 +52,11 @@ type StdNetBind struct { blackhole6 bool } -func NewStdNetBind() Bind { +func NewStdNetBind(externalControl control.Func) Bind { return &StdNetBind{ + externalControl: externalControl, + reservedForEndpoint: make(map[netip.AddrPort][3]uint8), + udpAddrPool: sync.Pool{ New: func() any { return &net.UDPAddr{ @@ -117,8 +126,29 @@ func (e *StdNetEndpoint) DstToString() string { return e.AddrPort.String() } -func listenNet(network string, port int) (*net.UDPConn, int, error) { - conn, err := listenConfig().ListenPacket(context.Background(), network, ":"+strconv.Itoa(port)) +func listenNet(externalControl control.Func, network string, port int) (*net.UDPConn, int, error) { + var listenerAddr string + if network == "udp6" { + listenerAddr = "[::]:" + strconv.Itoa(port) + } else { + listenerAddr = ":" + strconv.Itoa(port) + } + + var listener net.ListenConfig + listener.Control = func(network, address string, conn syscall.RawConn) error { + for _, wgControlFn := range controlFns { + err := wgControlFn(network, address, conn) + if err != nil { + return err + } + } + if externalControl != nil { + return externalControl(network, address, conn) + } else { + return nil + } + } + conn, err := listener.ListenPacket(context.Background(), network, listenerAddr) if err != nil { return nil, 0, err } @@ -154,13 +184,13 @@ again: var v4pc *ipv4.PacketConn var v6pc *ipv6.PacketConn - v4conn, port, err = listenNet("udp4", port) + v4conn, port, err = listenNet(s.externalControl, "udp4", port) if err != nil && !errors.Is(err, syscall.EAFNOSUPPORT) { return nil, 0, err } // Listen on the same port as we're using for ipv4. - v6conn, port, err = listenNet("udp6", port) + v6conn, port, err = listenNet(s.externalControl, "udp6", port) if uport == 0 && errors.Is(err, syscall.EADDRINUSE) && tries < 100 { v4conn.Close() tries++ @@ -265,8 +295,10 @@ func (s *StdNetBind) receiveIP( if sizes[i] == 0 { continue } - addrPort := msg.Addr.(*net.UDPAddr).AddrPort() - ep := &StdNetEndpoint{AddrPort: addrPort} // TODO: remove allocation + if msg.N > 3 { + common.ClearArray(bufs[i][1:4]) + } + ep := &StdNetEndpoint{AddrPort: M.AddrPortFromNet(msg.Addr)} // TODO: remove allocation getSrcFromControl(msg.OOB[:msg.NN], ep) eps[i] = ep } @@ -375,6 +407,14 @@ func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint, offset int) error { retried bool err error ) + for _, buf := range bufs { + if len(buf) > 3 { + reserved, loaded := s.reservedForEndpoint[endpoint.(*StdNetEndpoint).AddrPort] + if loaded { + copy(buf[1:4], reserved[:]) + } + } + } retry: if offload { n := coalesceMessages(ua, endpoint.(*StdNetEndpoint), bufs, offset, *msgs, setGSOSize) @@ -405,6 +445,10 @@ retry: return err } +func (s *StdNetBind) SetReservedForEndpoint(destination netip.AddrPort, reserved [3]byte) { + s.reservedForEndpoint[destination] = reserved +} + func (s *StdNetBind) send(conn *net.UDPConn, pc batchWriter, msgs []ipv6.Message) error { var ( n int diff --git a/conn/bind_windows.go b/conn/bind_windows.go index 66d91cd..121079f 100644 --- a/conn/bind_windows.go +++ b/conn/bind_windows.go @@ -15,6 +15,10 @@ import ( "sync/atomic" "unsafe" + "github.com/sagernet/sing/common" + "github.com/sagernet/sing/common/control" + E "github.com/sagernet/sing/common/exceptions" + M "github.com/sagernet/sing/common/metadata" "github.com/sagernet/wireguard-go/conn/winrio" "golang.org/x/sys/windows" ) @@ -71,18 +75,26 @@ type afWinRingBind struct { // WinRingBind uses Windows registered I/O for fast ring buffered networking. type WinRingBind struct { + externalControl control.Func + reservedForEndpoint map[WinRingEndpoint][3]uint8 + v4, v6 afWinRingBind mu sync.RWMutex isOpen atomic.Uint32 // 0, 1, or 2 } -func NewDefaultBind() Bind { return NewWinRingBind() } +func NewDefaultBind(externalControl control.Func) Bind { + return NewWinRingBind(externalControl) +} -func NewWinRingBind() Bind { +func NewWinRingBind(externalControl control.Func) Bind { if !winrio.Initialize() { - return NewStdNetBind() + return NewStdNetBind(externalControl) + } + return &WinRingBind{ + externalControl: externalControl, + reservedForEndpoint: make(map[WinRingEndpoint][3]uint8), } - return new(WinRingBind) } type WinRingEndpoint struct { @@ -238,7 +250,7 @@ func (ring *ringBuffer) Open() error { return nil } -func (bind *afWinRingBind) Open(family int32, sa windows.Sockaddr) (windows.Sockaddr, error) { +func (bind *afWinRingBind) Open(family int32, sa windows.Sockaddr, externalControl control.Func) (windows.Sockaddr, error) { var err error bind.sock, err = winrio.Socket(family, windows.SOCK_DGRAM, windows.IPPROTO_UDP) if err != nil { @@ -256,6 +268,19 @@ func (bind *afWinRingBind) Open(family int32, sa windows.Sockaddr) (windows.Sock if err != nil { return nil, err } + var network string + switch family { + case windows.AF_INET: + network = "udp4" + case windows.AF_INET6: + network = "udp6" + } + if externalControl != nil { + err = externalControl(network, M.AddrPortFromSockaddr(sa).String(), &fakeRawConn{bind.sock}) + if err != nil { + return nil, err + } + } err = windows.Bind(bind.sock, sa) if err != nil { return nil, err @@ -267,6 +292,23 @@ func (bind *afWinRingBind) Open(family int32, sa windows.Sockaddr) (windows.Sock return sa, nil } +type fakeRawConn struct { + socket windows.Handle +} + +func (c *fakeRawConn) Control(f func(fd uintptr)) error { + f(uintptr(c.socket)) + return nil +} + +func (c *fakeRawConn) Read(f func(fd uintptr) (done bool)) error { + panic("not implemented") +} + +func (c *fakeRawConn) Write(f func(fd uintptr) (done bool)) error { + panic("not implemented") +} + func (bind *WinRingBind) Open(port uint16) (recvFns []ReceiveFunc, selectedPort uint16, err error) { bind.mu.Lock() defer bind.mu.Unlock() @@ -279,11 +321,11 @@ func (bind *WinRingBind) Open(port uint16) (recvFns []ReceiveFunc, selectedPort return nil, 0, ErrBindAlreadyOpen } var sa windows.Sockaddr - sa, err = bind.v4.Open(windows.AF_INET, &windows.SockaddrInet4{Port: int(port)}) + sa, err = bind.v4.Open(windows.AF_INET, &windows.SockaddrInet4{Port: int(port)}, bind.externalControl) if err != nil { return nil, 0, err } - sa, err = bind.v6.Open(windows.AF_INET6, &windows.SockaddrInet6{Port: sa.(*windows.SockaddrInet4).Port}) + sa, err = bind.v6.Open(windows.AF_INET6, &windows.SockaddrInet6{Port: sa.(*windows.SockaddrInet4).Port}, bind.externalControl) if err != nil { return nil, 0, err } @@ -419,6 +461,9 @@ func (bind *WinRingBind) receiveIPv4(bufs [][]byte, sizes []int, eps []Endpoint) bind.mu.RLock() defer bind.mu.RUnlock() n, ep, err := bind.v4.Receive(bufs[0], &bind.isOpen) + if n > 3 { + common.ClearArray(bufs[0][1:4]) + } sizes[0] = n eps[0] = ep return 1, err @@ -428,6 +473,9 @@ func (bind *WinRingBind) receiveIPv6(bufs [][]byte, sizes []int, eps []Endpoint) bind.mu.RLock() defer bind.mu.RUnlock() n, ep, err := bind.v6.Receive(bufs[0], &bind.isOpen) + if n > 3 { + common.ClearArray(bufs[0][1:4]) + } sizes[0] = n eps[0] = ep return 1, err @@ -494,6 +542,12 @@ func (bind *WinRingBind) Send(bufs [][]byte, endpoint Endpoint, offset int) erro defer bind.mu.RUnlock() for _, buf := range bufs { buf = buf[offset:] + if len(buf) > 3 { + reserved, loaded := bind.reservedForEndpoint[*endpoint.(*WinRingEndpoint)] + if loaded { + copy(buf[1:4], reserved[:]) + } + } switch nend.family { case windows.AF_INET: if bind.v4.blackhole { @@ -514,6 +568,14 @@ func (bind *WinRingBind) Send(bufs [][]byte, endpoint Endpoint, offset int) erro return nil } +func (bind *WinRingBind) SetReservedForEndpoint(destination netip.AddrPort, reserved [3]byte) { + endpoint, err := bind.ParseEndpoint(destination.String()) + if err != nil { + panic(E.Cause(err, "parse destination as WinRingEndpoint")) + } + bind.reservedForEndpoint[*endpoint.(*WinRingEndpoint)] = reserved +} + func (s *StdNetBind) BindSocketToInterface4(interfaceIndex uint32, blackhole bool) error { s.mu.Lock() defer s.mu.Unlock() diff --git a/conn/conn.go b/conn/conn.go index 87b7510..c949b59 100644 --- a/conn/conn.go +++ b/conn/conn.go @@ -57,6 +57,8 @@ type Bind interface { // BatchSize is the number of buffers expected to be passed to // the ReceiveFuncs, and the maximum expected to be passed to SendBatch. BatchSize() int + + SetReservedForEndpoint(destination netip.AddrPort, reserved [3]byte) } // BindSocketToInterface is implemented by Bind objects that support being diff --git a/conn/controlfns.go b/conn/controlfns.go index 27421bd..d4164e4 100644 --- a/conn/controlfns.go +++ b/conn/controlfns.go @@ -6,8 +6,7 @@ package conn import ( - "net" - "syscall" + "github.com/sagernet/sing/common/control" ) // UDP socket read/write buffer size (7MB). The value of 7MB is chosen as it is @@ -17,27 +16,6 @@ import ( // around this limitation) const socketBufferSize = 7 << 20 -// controlFn is the callback function signature from net.ListenConfig.Control. -// It is used to apply platform specific configuration to the socket prior to -// bind. -type controlFn func(network, address string, c syscall.RawConn) error - // controlFns is a list of functions that are called from the listen config // that can apply socket options. -var controlFns = []controlFn{} - -// listenConfig returns a net.ListenConfig that applies the controlFns to the -// socket prior to bind. This is used to apply socket buffer sizing and packet -// information OOB configuration for sticky sockets. -func listenConfig() *net.ListenConfig { - return &net.ListenConfig{ - Control: func(network, address string, c syscall.RawConn) error { - for _, fn := range controlFns { - if err := fn(network, address, c); err != nil { - return err - } - } - return nil - }, - } -} +var controlFns []control.Func diff --git a/conn/default.go b/conn/default.go index 2ce1579..9907f64 100644 --- a/conn/default.go +++ b/conn/default.go @@ -7,4 +7,8 @@ package conn -func NewDefaultBind() Bind { return NewStdNetBind() } +import "github.com/sagernet/sing/common/control" + +func NewDefaultBind(externalControl control.Func) Bind { + return NewStdNetBind(externalControl) +} From fa73d0f1ae147fd4a214b20e018a8e2c6c4ba6b2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Wed, 17 Sep 2025 19:03:11 +0800 Subject: [PATCH 32/34] Fix input packet --- device/send.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/device/send.go b/device/send.go index 64a4221..264e964 100644 --- a/device/send.go +++ b/device/send.go @@ -328,7 +328,7 @@ func (device *Device) InputPacket(destination []byte, packetSlices [][]byte) { return } elem := device.NewOutboundElement() - packet := elem.buffer[MessageTransportHeaderSize:] + packet := elem.buffer[MessageEncapsulatingTransportSize+MessageTransportHeaderSize:] var n int for _, packetSlice := range packetSlices { n += copy(packet[n:], packetSlice) From 19b0d3587703f82379d4d05b03e8666d16df1b7a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Tue, 24 Feb 2026 15:47:47 +0800 Subject: [PATCH 33/34] Fix reserved bytes offset in StdNetBind.Send --- conn/bind_std.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/conn/bind_std.go b/conn/bind_std.go index cb979e3..4dadb0d 100644 --- a/conn/bind_std.go +++ b/conn/bind_std.go @@ -408,10 +408,10 @@ func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint, offset int) error { err error ) for _, buf := range bufs { - if len(buf) > 3 { + if len(buf) > offset+3 { reserved, loaded := s.reservedForEndpoint[endpoint.(*StdNetEndpoint).AddrPort] if loaded { - copy(buf[1:4], reserved[:]) + copy(buf[offset+1:offset+4], reserved[:]) } } } From e5feca7d61f3e88fd34b7b8bda377d35fbd4157e Mon Sep 17 00:00:00 2001 From: Leadaxe <247031499+Leadaxe@users.noreply.github.com> Date: Tue, 23 Jun 2026 02:31:21 +0300 Subject: [PATCH 34/34] lx: re-graft AmneziaWG 2.0 obfuscation onto sagernet/wireguard-go v0.0.3 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Migrate the AmneziaWG graft from the old sagernet base (506b763) onto v0.0.3, the base sing-box 1.14 uses. Strategy: start from v0.0.3, copy the net-new obfuscation files verbatim from the proven graft (27290b6), take the 6 modified device/ files wholesale from the graft, keep all tun/* and conn/* from v0.0.3. Net-new (copied from 27290b6): device/magic-header.go + device/obf*.go (9). Modified (from 27290b6): device/{send,receive,uapi,device,noise-protocol,cookie}.go. - receive.go: reverted two Go-1.22 'range int' loops to the v0.0.3 form (go 1.20). - cookie.go: keeps the 4-arg CreateReply(msgType) the grafted send.go calls. - noise-protocol.go: MessageEncapsulatingTransportSize=0 (AWG composes without the sagernet Bind.Send headroom prepend). §010 android UDP_GRO guard (fb8d8d8) intentionally DROPPED: v0.0.3 fixes the receive split-brain at source — bind_std.go now gates all 5 RX/offload paths on linux||android, and gso_linux.go splits coalesced packets on android. Re-inserting the guard would disable a now-working android RX-offload path. Requires on-device re-verification before release. go.mod: go 1.20, golang.org/x/sys v0.21.0 (v0.0.3 baseline). No amnezia-vpn import paths remain. Builds clean for linux/android/windows/darwin (library packages). --- device/cookie.go | 3 +- device/device.go | 31 ++++- device/magic-header.go | 63 +++++++++ device/noise-protocol.go | 19 +-- device/obf.go | 140 +++++++++++++++++++ device/obf_bytes.go | 47 +++++++ device/obf_data.go | 25 ++++ device/obf_datasize.go | 38 +++++ device/obf_datastring.go | 29 ++++ device/obf_rand.go | 39 ++++++ device/obf_randchars.go | 48 +++++++ device/obf_randdigits.go | 48 +++++++ device/obf_timestamp.go | 31 +++++ device/receive.go | 107 ++++++++++++-- device/send.go | 107 +++++++++++--- device/uapi.go | 291 +++++++++++++++++++++++++++++++++++++-- 16 files changed, 1018 insertions(+), 48 deletions(-) create mode 100644 device/magic-header.go create mode 100644 device/obf.go create mode 100644 device/obf_bytes.go create mode 100644 device/obf_data.go create mode 100644 device/obf_datasize.go create mode 100644 device/obf_datastring.go create mode 100644 device/obf_rand.go create mode 100644 device/obf_randchars.go create mode 100644 device/obf_randdigits.go create mode 100644 device/obf_timestamp.go diff --git a/device/cookie.go b/device/cookie.go index a093c8b..6a0463c 100644 --- a/device/cookie.go +++ b/device/cookie.go @@ -118,6 +118,7 @@ func (st *CookieChecker) CreateReply( msg []byte, recv uint32, src []byte, + msgType uint32, ) (*MessageCookieReply, error) { st.RLock() @@ -153,7 +154,7 @@ func (st *CookieChecker) CreateReply( smac1 := smac2 - blake2s.Size128 reply := new(MessageCookieReply) - reply.Type = MessageCookieReplyType + reply.Type = msgType reply.Receiver = recv _, err := rand.Read(reply.Nonce[:]) diff --git a/device/device.go b/device/device.go index 8afa6fd..140a5df 100644 --- a/device/device.go +++ b/device/device.go @@ -93,6 +93,29 @@ type Device struct { closed chan struct{} log *Logger pauseManager pause.Manager + + // lx: AmneziaWG obfuscation state (grafted from amneziawg-go). + junk struct { + min int + max int + count int + } + + headers struct { + init *magicHeader + cookie *magicHeader + response *magicHeader + transport *magicHeader + } + + paddings struct { + init int + response int + cookie int + transport int + } + + ipackets [5]*obfChain } // deviceState represents the state of a Device. @@ -166,7 +189,8 @@ func (device *Device) changeState(want deviceState) (err error) { err = errDown } } - device.log.Verbosef("Interface state was %s, requested %s, now %s", old, want, device.deviceState()) + device.log.Verbosef( + "Interface state was %s, requested %s, now %s", old, want, device.deviceState()) return } @@ -303,6 +327,11 @@ func NewDevice(ctx context.Context, tunDevice tun.Device, bind conn.Bind, logger device.rate.limiter.Init() device.indexTable.Init() + device.headers.init = &magicHeader{start: MessageInitiationType, end: MessageInitiationType} + device.headers.response = &magicHeader{start: MessageResponseType, end: MessageResponseType} + device.headers.cookie = &magicHeader{start: MessageCookieReplyType, end: MessageCookieReplyType} + device.headers.transport = &magicHeader{start: MessageTransportType, end: MessageTransportType} + device.PopulatePools() // create queues diff --git a/device/magic-header.go b/device/magic-header.go new file mode 100644 index 0000000..78e59d6 --- /dev/null +++ b/device/magic-header.go @@ -0,0 +1,63 @@ +package device + +import ( + "crypto/rand" + "errors" + "fmt" + "math/big" + "strconv" + "strings" +) + +type magicHeader struct { + start uint32 + end uint32 +} + +func newMagicHeader(spec string) (*magicHeader, error) { + parts := strings.Split(spec, "-") + if len(parts) < 1 || len(parts) > 2 { + return nil, errors.New("bad format") + } + + start, err := strconv.ParseUint(parts[0], 10, 32) + if err != nil { + return nil, fmt.Errorf("failed to parse %s: %w", parts[0], err) + } + + var end uint64 + if len(parts) > 1 { + end, err = strconv.ParseUint(parts[1], 10, 32) + if err != nil { + return nil, fmt.Errorf("failed to parse %s: %w", parts[1], err) + } + } else { + end = start + } + + if end < start { + return nil, errors.New("wrong range specified") + } + + return &magicHeader{ + start: uint32(start), + end: uint32(end), + }, nil +} + +func (h *magicHeader) GenSpec() string { + if h.start == h.end { + return fmt.Sprintf("%d", h.start) + } + return fmt.Sprintf("%d-%d", h.start, h.end) +} + +func (h *magicHeader) Validate(val uint32) bool { + return h.start <= val && val <= h.end +} + +func (h *magicHeader) Generate() uint32 { + high := int64(h.end - h.start + 1) + r, _ := rand.Int(rand.Reader, big.NewInt(high)) + return h.start + uint32(r.Int64()) +} diff --git a/device/noise-protocol.go b/device/noise-protocol.go index ed4d82a..745edd1 100644 --- a/device/noise-protocol.go +++ b/device/noise-protocol.go @@ -54,10 +54,11 @@ const ( ) const ( - MessageInitiationType = 1 - MessageResponseType = 2 - MessageCookieReplyType = 3 - MessageTransportType = 4 + MessageUnknownType uint32 = 0 + MessageInitiationType uint32 = 1 + MessageResponseType uint32 = 2 + MessageCookieReplyType uint32 = 3 + MessageTransportType uint32 = 4 ) const ( @@ -65,7 +66,7 @@ const ( MessageResponseSize = 92 // size of response message MessageCookieReplySize = 64 // size of cookie reply message MessageTransportHeaderSize = 16 // size of data preceding content in transport message - MessageEncapsulatingTransportSize = 8 // size of optional, free (for use by conn.Bind.Send()) space preceding the transport header + MessageEncapsulatingTransportSize = 0 // lx: zeroed so AmneziaWG obfuscation composes without sagernet headroom (AWG path doesn't use the Bind.Send prepend) MessageTransportSize = MessageTransportHeaderSize + poly1305.TagSize // size of empty transport MessageKeepaliveSize = MessageTransportSize // size of keepalive MessageHandshakeSize = MessageInitiationSize // size of largest handshake related message @@ -218,7 +219,7 @@ type Handshake struct { localEphemeral NoisePrivateKey // ephemeral secret key localIndex uint32 // used to clear hash-table remoteIndex uint32 // index for sending - remoteStatic NoisePublicKey // long term key + remoteStatic NoisePublicKey // long term key, never changes, can be accessed without mutex remoteEphemeral NoisePublicKey // ephemeral public key precomputedStaticStatic [NoisePublicKeySize]byte // precomputed shared secret lastTimestamp tai64n.Timestamp @@ -287,8 +288,10 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e handshake.mixHash(handshake.remoteStatic[:]) + msgType := device.headers.init.Generate() + msg := MessageInitiation{ - Type: MessageInitiationType, + Type: msgType, Ephemeral: handshake.localEphemeral.publicKey(), } @@ -466,7 +469,7 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error } var msg MessageResponse - msg.Type = MessageResponseType + msg.Type = device.headers.response.Generate() msg.Sender = handshake.localIndex msg.Receiver = handshake.remoteIndex diff --git a/device/obf.go b/device/obf.go new file mode 100644 index 0000000..53c55ff --- /dev/null +++ b/device/obf.go @@ -0,0 +1,140 @@ +package device + +import ( + "errors" + "fmt" + "strings" +) + +type obfBuilder func(val string) (obf, error) + +var obfBuilders = map[string]obfBuilder{ + "b": newBytesObf, + "t": newTimestampObf, + "r": newRandObf, + "rc": newRandCharObf, + "rd": newRandDigitsObf, + "d": newDataObf, + "ds": newDataStringObf, + "dz": newDataSizeObf, +} + +type obf interface { + Obfuscate(dst, src []byte) + Deobfuscate(dst, src []byte) bool + ObfuscatedLen(srcLen int) int + DeobfuscatedLen(srcLen int) int +} + +type obfChain struct { + Spec string + obfs []obf +} + +func newObfChain(spec string) (*obfChain, error) { + var ( + obfs []obf + errs []error + ) + + remaining := spec[:] + for { + start := strings.IndexByte(remaining, '<') + if start == -1 { + break + } + + end := strings.IndexByte(remaining[start:], '>') + if end == -1 { + return nil, errors.New("missing enclosing >") + } + end += start + + tag := remaining[start+1 : end] + parts := strings.Fields(tag) + if len(parts) == 0 { + errs = append(errs, errors.New("empty tag")) + remaining = remaining[end+1:] + continue + } + + key := parts[0] + builder, ok := obfBuilders[key] + if !ok { + errs = append(errs, fmt.Errorf("unknown tag <%s>", key)) + remaining = remaining[end+1:] + continue + } + + val := "" + if len(parts) > 1 { + val = parts[1] + } + + o, err := builder(val) + if err != nil { + errs = append(errs, fmt.Errorf("failed to build <%s>: %w", key, err)) + remaining = remaining[end+1:] + continue + } + + obfs = append(obfs, o) + remaining = remaining[end+1:] + } + + if len(errs) > 0 { + return nil, errors.Join(errs...) + } + + return &obfChain{ + Spec: spec, + obfs: obfs, + }, nil +} + +func (c *obfChain) Obfuscate(dst, src []byte) { + written := 0 + for _, o := range c.obfs { + obfLen := o.ObfuscatedLen(len(src)) + o.Obfuscate(dst[written:written+obfLen], src) + written += obfLen + } +} + +func (c *obfChain) Deobfuscate(dst, src []byte) bool { + dynamicLen := len(src) - c.ObfuscatedLen(0) + + written, read := 0, 0 + + for _, o := range c.obfs { + deobfLen := o.DeobfuscatedLen(dynamicLen) + obfLen := o.ObfuscatedLen(deobfLen) + + if !o.Deobfuscate(dst[written:written+deobfLen], src[read:read+obfLen]) { + return false + } + + written += deobfLen + read += obfLen + } + + return true +} + +func (c *obfChain) ObfuscatedLen(n int) int { + total := 0 + for _, o := range c.obfs { + total += o.ObfuscatedLen(n) + } + return total +} + +func (c *obfChain) DeobfuscatedLen(n int) int { + dynamicLen := n - c.ObfuscatedLen(0) + + total := 0 + for _, o := range c.obfs { + total += o.DeobfuscatedLen(dynamicLen) + } + return total +} diff --git a/device/obf_bytes.go b/device/obf_bytes.go new file mode 100644 index 0000000..68d722b --- /dev/null +++ b/device/obf_bytes.go @@ -0,0 +1,47 @@ +package device + +import ( + "bytes" + "encoding/hex" + "errors" + "strings" +) + +func newBytesObf(val string) (obf, error) { + val = strings.TrimPrefix(val, "0x") + + if len(val) == 0 { + return nil, errors.New("empty argument") + } + + if len(val)%2 != 0 { + return nil, errors.New("odd amount of symbols") + } + + bytes, err := hex.DecodeString(val) + if err != nil { + return nil, err + } + + return &bytesObf{data: bytes}, nil +} + +type bytesObf struct { + data []byte +} + +func (o *bytesObf) Obfuscate(dst, src []byte) { + copy(dst, o.data) +} + +func (o *bytesObf) Deobfuscate(dst, src []byte) bool { + return bytes.Equal(o.data, src[:o.ObfuscatedLen(0)]) +} + +func (o *bytesObf) ObfuscatedLen(srcLen int) int { + return len(o.data) +} + +func (o *bytesObf) DeobfuscatedLen(srcLen int) int { + return 0 +} diff --git a/device/obf_data.go b/device/obf_data.go new file mode 100644 index 0000000..42d3f65 --- /dev/null +++ b/device/obf_data.go @@ -0,0 +1,25 @@ +package device + +func newDataObf(val string) (obf, error) { + return &dataObf{}, nil +} + +type dataObf struct { +} + +func (obf *dataObf) Obfuscate(dst, src []byte) { + copy(dst, src) +} + +func (obf *dataObf) Deobfuscate(dst, src []byte) bool { + copy(dst, src) + return true +} + +func (o *dataObf) ObfuscatedLen(n int) int { + return n +} + +func (o *dataObf) DeobfuscatedLen(n int) int { + return n +} diff --git a/device/obf_datasize.go b/device/obf_datasize.go new file mode 100644 index 0000000..7267e2a --- /dev/null +++ b/device/obf_datasize.go @@ -0,0 +1,38 @@ +package device + +import "strconv" + +func newDataSizeObf(val string) (obf, error) { + length, err := strconv.Atoi(val) + if err != nil { + return nil, err + } + + return &dataSizeObf{ + length: length, + }, nil +} + +type dataSizeObf struct { + length int +} + +func (o *dataSizeObf) Obfuscate(dst, src []byte) { + srcLen := len(src) + for i := o.length - 1; i >= 0; i-- { + dst[i] = byte(srcLen & 0xFF) + srcLen >>= 8 + } +} + +func (o *dataSizeObf) Deobfuscate(dst, src []byte) bool { + return true +} + +func (o *dataSizeObf) ObfuscatedLen(srcLen int) int { + return o.length +} + +func (o *dataSizeObf) DeobfuscatedLen(srcLen int) int { + return 0 +} diff --git a/device/obf_datastring.go b/device/obf_datastring.go new file mode 100644 index 0000000..2701e95 --- /dev/null +++ b/device/obf_datastring.go @@ -0,0 +1,29 @@ +package device + +import ( + "encoding/base64" +) + +func newDataStringObf(val string) (obf, error) { + return &dataStringObf{}, nil +} + +type dataStringObf struct { +} + +func (o *dataStringObf) Obfuscate(dst, src []byte) { + base64.RawStdEncoding.Encode(dst, src) +} + +func (o *dataStringObf) Deobfuscate(dst, src []byte) bool { + base64.RawStdEncoding.Decode(dst, src) + return true +} + +func (o *dataStringObf) ObfuscatedLen(n int) int { + return base64.RawStdEncoding.EncodedLen(n) +} + +func (o *dataStringObf) DeobfuscatedLen(n int) int { + return base64.RawStdEncoding.DecodedLen(n) +} diff --git a/device/obf_rand.go b/device/obf_rand.go new file mode 100644 index 0000000..edf461e --- /dev/null +++ b/device/obf_rand.go @@ -0,0 +1,39 @@ +package device + +import ( + "crypto/rand" + "strconv" +) + +func newRandObf(val string) (obf, error) { + length, err := strconv.Atoi(val) + if err != nil { + return nil, err + } + + return &randObf{ + length: length, + }, nil +} + +type randObf struct { + length int +} + +func (o *randObf) Obfuscate(dst, src []byte) { + rand.Read(dst[:o.length]) +} + +func (o *randObf) Deobfuscate(dst, src []byte) bool { + // there is no way to validate randomness :) + // assume that it is always true + return true +} + +func (o *randObf) ObfuscatedLen(n int) int { + return o.length +} + +func (o *randObf) DeobfuscatedLen(n int) int { + return 0 +} diff --git a/device/obf_randchars.go b/device/obf_randchars.go new file mode 100644 index 0000000..1d9968c --- /dev/null +++ b/device/obf_randchars.go @@ -0,0 +1,48 @@ +package device + +import ( + "crypto/rand" + "strconv" + "unicode" +) + +const chars52 = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ" + +func newRandCharObf(val string) (obf, error) { + length, err := strconv.Atoi(val) + if err != nil { + return nil, err + } + + return &randCharObf{ + length: length, + }, nil +} + +type randCharObf struct { + length int +} + +func (o *randCharObf) Obfuscate(dst, src []byte) { + rand.Read(dst[:o.length]) + for i := range dst[:o.length] { + dst[i] = chars52[dst[i]%52] + } +} + +func (o *randCharObf) Deobfuscate(dst, src []byte) bool { + for _, b := range src[:o.length] { + if !unicode.IsLetter(rune(b)) { + return false + } + } + return true +} + +func (o *randCharObf) ObfuscatedLen(n int) int { + return o.length +} + +func (o *randCharObf) DeobfuscatedLen(n int) int { + return 0 +} diff --git a/device/obf_randdigits.go b/device/obf_randdigits.go new file mode 100644 index 0000000..4794bb1 --- /dev/null +++ b/device/obf_randdigits.go @@ -0,0 +1,48 @@ +package device + +import ( + "crypto/rand" + "strconv" + "unicode" +) + +const digits10 = "0123456789" + +func newRandDigitsObf(val string) (obf, error) { + length, err := strconv.Atoi(val) + if err != nil { + return nil, err + } + + return &randDigitObf{ + length: length, + }, nil +} + +type randDigitObf struct { + length int +} + +func (o *randDigitObf) Obfuscate(dst, src []byte) { + rand.Read(dst[:o.length]) + for i := range dst[:o.length] { + dst[i] = digits10[dst[i]%10] + } +} + +func (o *randDigitObf) Deobfuscate(dst, src []byte) bool { + for _, b := range src[:o.length] { + if !unicode.IsDigit(rune(b)) { + return false + } + } + return true +} + +func (o *randDigitObf) ObfuscatedLen(n int) int { + return o.length +} + +func (o *randDigitObf) DeobfuscatedLen(n int) int { + return 0 +} diff --git a/device/obf_timestamp.go b/device/obf_timestamp.go new file mode 100644 index 0000000..0a8180b --- /dev/null +++ b/device/obf_timestamp.go @@ -0,0 +1,31 @@ +package device + +import ( + "encoding/binary" + "time" +) + +func newTimestampObf(_ string) (obf, error) { + return ×tampObf{}, nil +} + +type timestampObf struct{} + +func (o *timestampObf) Obfuscate(dst, src []byte) { + t := uint32(time.Now().Unix()) + binary.BigEndian.PutUint32(dst, t) +} + +func (o *timestampObf) Deobfuscate(dst, src []byte) bool { + // replay attack check? + // requires time to be always synchronized + return true +} + +func (o *timestampObf) ObfuscatedLen(n int) int { + return 4 +} + +func (o *timestampObf) DeobfuscatedLen(n int) int { + return 0 +} diff --git a/device/receive.go b/device/receive.go index e13c987..93117c4 100644 --- a/device/receive.go +++ b/device/receive.go @@ -69,7 +69,10 @@ func (peer *Peer) keepKeyFreshReceiving() { * Every time the bind is updated a new routine is started for * IPv4 and IPv6 (separately) */ -func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.ReceiveFunc) { +func (device *Device) RoutineReceiveIncoming( + maxBatchSize int, + recv conn.ReceiveFunc, +) { recvName := recv.PrettyName() defer func() { device.log.Verbosef("Routine: receive incoming %s - stopped", recvName) @@ -132,9 +135,14 @@ func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.Receive } // check size of packet - packet := bufsArrs[i][:size] - msgType := binary.LittleEndian.Uint32(packet[:4]) + + // get message padding and type based on information from S1-S4 and H1-H4 + msgType, padding := device.DeterminePacketTypeAndPadding(packet, MessageUnknownType) + if padding > 0 { + copy(packet, packet[padding:]) + packet = packet[:len(packet)-padding] + } switch msgType { @@ -276,7 +284,6 @@ func (device *Device) RoutineHandshake(id int) { device.log.Verbosef("Routine: handshake worker %d - started", id) for elem := range device.queue.handshake.c { - // handle cookie fields and ratelimiting switch elem.msgType { @@ -303,9 +310,14 @@ func (device *Device) RoutineHandshake(id int) { // consume reply if peer := entry.peer; peer.isRunning.Load() { - device.log.Verbosef("Receiving cookie response from %s", elem.endpoint.DstToString()) + device.log.Verbosef( + "Receiving cookie response from %s", + elem.endpoint.DstToString(), + ) if !peer.cookieGenerator.ConsumeReply(&reply) { - device.log.Verbosef("Could not decrypt invalid cookie response") + device.log.Verbosef( + "Could not decrypt invalid cookie response", + ) } } @@ -347,9 +359,7 @@ func (device *Device) RoutineHandshake(id int) { switch elem.msgType { case MessageInitiationType: - // unmarshal - var msg MessageInitiation err := msg.unmarshal(elem.packet) if err != nil { @@ -357,7 +367,8 @@ func (device *Device) RoutineHandshake(id int) { goto skip } - // consume initiation + // have to reassign msgType for ranged msgType to work + msg.Type = elem.msgType peer := device.ConsumeMessageInitiation(&msg, elem.endpoint) if peer == nil { @@ -389,6 +400,9 @@ func (device *Device) RoutineHandshake(id int) { goto skip } + // have to reassign msgType for ranged msgType to work + msg.Type = elem.msgType + // consume response peer := device.ConsumeMessageResponse(&msg) @@ -505,11 +519,28 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) { } default: - device.log.Verbosef("Packet with invalid IP version from %v", peer) + device.log.Verbosef( + "Packet with invalid IP version from %v", + peer, + ) continue } - bufs = append(bufs, elem.buffer[:MessageTransportOffsetContent+len(elem.packet)]) + bufs = append( + bufs, + elem.buffer[:MessageTransportOffsetContent+len(elem.packet)], + ) + } + + peer.rxBytes.Add(rxBytesLen) + if validTailPacket >= 0 { + peer.SetEndpointFromPacket(elemsContainer.elems[validTailPacket].endpoint) + peer.keepKeyFreshReceiving() + peer.timersAnyAuthenticatedPacketTraversal() + peer.timersAnyAuthenticatedPacketReceived() + } + if dataPacketReceived { + peer.timersDataReceived() } peer.rxBytes.Add(rxBytesLen) @@ -536,3 +567,57 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) { device.PutInboundElementsContainer(elemsContainer) } } + +func (device *Device) DeterminePacketTypeAndPadding(packet []byte, expectedType uint32) (uint32, int) { + size := len(packet) + + if expectedType == MessageUnknownType || expectedType == MessageInitiationType { + padding := device.paddings.init + header := device.headers.init + + if size == padding+MessageInitiationSize { + data := packet[padding:] + if header.Validate(binary.LittleEndian.Uint32(data)) { + return MessageInitiationType, padding + } + } + } + + if expectedType == MessageUnknownType || expectedType == MessageResponseType { + padding := device.paddings.response + header := device.headers.response + + if size == padding+MessageResponseSize { + data := packet[padding:] + if header.Validate(binary.LittleEndian.Uint32(data)) { + return MessageResponseType, padding + } + } + } + + if expectedType == MessageUnknownType || expectedType == MessageCookieReplyType { + padding := device.paddings.cookie + header := device.headers.cookie + + if size == padding+MessageCookieReplySize { + data := packet[padding:] + if header.Validate(binary.LittleEndian.Uint32(data)) { + return MessageCookieReplyType, padding + } + } + } + + if expectedType == MessageUnknownType || expectedType == MessageTransportType { + padding := device.paddings.transport + header := device.headers.transport + + if size >= padding+MessageTransportHeaderSize { + data := packet[padding:] + if header.Validate(binary.LittleEndian.Uint32(data)) { + return MessageTransportType, padding + } + } + } + + return MessageUnknownType, 0 +} diff --git a/device/send.go b/device/send.go index 264e964..7653f62 100644 --- a/device/send.go +++ b/device/send.go @@ -6,8 +6,11 @@ package device import ( + "bytes" + "crypto/rand" "encoding/binary" "errors" + "math/big" "net" "os" "sync" @@ -127,15 +130,48 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error { return err } - buf := make([]byte, MessageEncapsulatingTransportSize+MessageInitiationSize) - packet := buf[MessageEncapsulatingTransportSize:] - _ = msg.marshal(packet) + var sendBuffer [][]byte + + for _, ipacket := range peer.device.ipackets { + if ipacket != nil { + buf := make([]byte, ipacket.ObfuscatedLen(0)) + ipacket.Obfuscate(buf, nil) + sendBuffer = append(sendBuffer, buf) + } + } + + jc := peer.device.junk.count + jmin := peer.device.junk.min + jmax := peer.device.junk.max + + for i := 0; i < jc; i++ { + nBig, _ := rand.Int(rand.Reader, big.NewInt(int64(jmax-jmin+1))) + n := int(nBig.Int64()) + jmin + + buf := make([]byte, n) + rand.Read(buf) + sendBuffer = append(sendBuffer, buf) + } + + var buf [MessageInitiationSize]byte + writer := bytes.NewBuffer(buf[:0]) + binary.Write(writer, binary.LittleEndian, msg) + packet := writer.Bytes() peer.cookieGenerator.AddMacs(packet) peer.timersAnyAuthenticatedPacketTraversal() peer.timersAnyAuthenticatedPacketSent() - err = peer.SendBuffers([][]byte{buf}) + if padding := peer.device.paddings.init; padding > 0 { + buf := make([]byte, padding+len(packet)) + rand.Read(buf[:padding]) + copy(buf[padding:], packet) + packet = buf + } + + sendBuffer = append(sendBuffer, packet) + + err = peer.SendBuffers(sendBuffer) if err != nil { peer.device.log.Errorf("%v - Failed to send handshake initiation: %v", peer, err) } @@ -157,9 +193,11 @@ func (peer *Peer) SendHandshakeResponse() error { return err } - buf := make([]byte, MessageEncapsulatingTransportSize+MessageResponseSize) - packet := buf[MessageEncapsulatingTransportSize:] - _ = response.marshal(packet) + var buf [MessageResponseSize]byte + writer := bytes.NewBuffer(buf[:0]) + + binary.Write(writer, binary.LittleEndian, response) + packet := writer.Bytes() peer.cookieGenerator.AddMacs(packet) err = peer.BeginSymmetricSession() @@ -172,8 +210,15 @@ func (peer *Peer) SendHandshakeResponse() error { peer.timersAnyAuthenticatedPacketTraversal() peer.timersAnyAuthenticatedPacketSent() + if padding := peer.device.paddings.response; padding > 0 { + buf := make([]byte, padding+len(packet)) + rand.Read(buf[:padding]) + copy(buf[padding:], packet) + packet = buf + } + // TODO: allocation could be avoided - err = peer.SendBuffers([][]byte{buf}) + err = peer.SendBuffers([][]byte{packet}) if err != nil { peer.device.log.Errorf("%v - Failed to send handshake response: %v", peer, err) } @@ -184,18 +229,33 @@ func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement) device.log.Verbosef("Sending cookie response for denied handshake message for %v", initiatingElem.endpoint.DstToString()) sender := binary.LittleEndian.Uint32(initiatingElem.packet[4:8]) - reply, err := device.cookieChecker.CreateReply(initiatingElem.packet, sender, initiatingElem.endpoint.DstToBytes()) + msgType := device.headers.cookie.Generate() + + reply, err := device.cookieChecker.CreateReply( + initiatingElem.packet, + sender, + initiatingElem.endpoint.DstToBytes(), + msgType, + ) if err != nil { device.log.Errorf("Failed to create cookie reply: %v", err) return err } - buf := make([]byte, MessageEncapsulatingTransportSize+MessageCookieReplySize) - packet := buf[MessageEncapsulatingTransportSize:] - _ = reply.marshal(packet) - // TODO: allocation could be avoided - device.net.bind.Send([][]byte{buf}, initiatingElem.endpoint, MessageEncapsulatingTransportSize) + var buf [MessageCookieReplySize]byte + writer := bytes.NewBuffer(buf[:0]) + binary.Write(writer, binary.LittleEndian, reply) + packet := writer.Bytes() + if padding := device.paddings.cookie; padding > 0 { + buf := make([]byte, padding+len(packet)) + rand.Read(buf[:padding]) + copy(buf[padding:], packet) + packet = buf + } + + // TODO: allocation could be avoided + device.net.bind.Send([][]byte{packet}, initiatingElem.endpoint, 0) return nil } @@ -477,13 +537,15 @@ func (device *Device) RoutineEncryption(id int) { for elemsContainer := range device.queue.encryption.c { for _, elem := range elemsContainer.elems { // populate header fields - header := elem.buffer[MessageEncapsulatingTransportSize : MessageEncapsulatingTransportSize+MessageTransportHeaderSize] + header := elem.buffer[:MessageTransportHeaderSize] fieldType := header[0:4] fieldReceiver := header[4:8] fieldNonce := header[8:16] - binary.LittleEndian.PutUint32(fieldType, MessageTransportType) + msgType := device.headers.transport.Generate() + + binary.LittleEndian.PutUint32(fieldType, msgType) binary.LittleEndian.PutUint32(fieldReceiver, elem.keypair.remoteIndex) binary.LittleEndian.PutUint64(fieldNonce, elem.nonce) @@ -500,9 +562,6 @@ func (device *Device) RoutineEncryption(id int) { elem.packet, nil, ) - - // re-slice packet to include encapsulating transport space - elem.packet = elem.buffer[:MessageEncapsulatingTransportSize+len(elem.packet)] } elemsContainer.Unlock() } @@ -544,6 +603,15 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) { if len(elem.packet) != MessageKeepaliveSize { dataSent = true } + if padding := device.paddings.transport; padding > 0 { + // elem.packet is stored at the start of elem.buffer + // with zero padding + for i := len(elem.packet) - 1; i >= 0; i-- { + elem.buffer[i+padding] = elem.buffer[i] + } + rand.Read(elem.buffer[:padding]) + elem.packet = elem.buffer[:padding+len(elem.packet)] + } bufs = append(bufs, elem.packet) } @@ -554,6 +622,7 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) { if dataSent { peer.timersDataSent() } + for _, elem := range elemsContainer.elems { device.PutMessageBuffer(elem.buffer) device.PutOutboundElement(elem) diff --git a/device/uapi.go b/device/uapi.go index cba371d..4f295c1 100644 --- a/device/uapi.go +++ b/device/uapi.go @@ -97,6 +97,56 @@ func (device *Device) IpcGetOperation(w io.Writer) error { sendf("fwmark=%d", device.net.fwmark) } + if device.junk.count != 0 { + sendf("jc=%d", device.junk.count) + } + + if device.junk.min != 0 { + sendf("jmin=%d", device.junk.min) + } + + if device.junk.max != 0 { + sendf("jmax=%d", device.junk.max) + } + + if device.paddings.init != 0 { + sendf("s1=%d", device.paddings.init) + } + + if device.paddings.response != 0 { + sendf("s2=%d", device.paddings.response) + } + + if device.paddings.cookie != 0 { + sendf("s3=%d", device.paddings.cookie) + } + + if device.paddings.transport != 0 { + sendf("s4=%d", device.paddings.transport) + } + + if device.headers.init != nil { + sendf("h1=%s", device.headers.init.GenSpec()) + } + + if device.headers.response != nil { + sendf("h2=%s", device.headers.response.GenSpec()) + } + + if device.headers.cookie != nil { + sendf("h3=%s", device.headers.cookie.GenSpec()) + } + + if device.headers.transport != nil { + sendf("h4=%s", device.headers.transport.GenSpec()) + } + + for i, ipacket := range device.ipackets { + if ipacket != nil { + sendf("i%d=%s", i+1, ipacket.Spec) + } + } + for _, peer := range device.peers.keyMap { // Serialize peer state. peer.handshake.mutex.RLock() @@ -147,6 +197,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) { } }() + ipcDev := new(ipcSetDevice) peer := new(ipcSetPeer) deviceConfig := true @@ -155,12 +206,20 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) { line := scanner.Text() if line == "" { // Blank line means terminate operation. + err := ipcDev.mergeWithDevice(device) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to merge with device: %w", err) + } peer.handlePostConfig() return nil } key, value, ok := strings.Cut(line, "=") if !ok { - return ipcErrorf(ipc.IpcErrorProtocol, "failed to parse line %q", line) + return ipcErrorf( + ipc.IpcErrorProtocol, + "failed to parse line %q", + line, + ) } if key == "public_key" { @@ -186,6 +245,10 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) { return err } } + err = ipcDev.mergeWithDevice(device) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to merge with device: %w", err) + } peer.handlePostConfig() if err := scanner.Err(); err != nil { @@ -235,11 +298,155 @@ func (device *Device) handleDeviceLine(key, value string) error { case "replace_peers": if value != "true" { - return ipcErrorf(ipc.IpcErrorInvalid, "failed to set replace_peers, invalid value: %v", value) + return ipcErrorf( + ipc.IpcErrorInvalid, + "failed to set replace_peers, invalid value: %v", + value, + ) } device.log.Verbosef("UAPI: Removing all peers") device.RemoveAllPeers() + case "jc": + jc, err := strconv.Atoi(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jc: %w", err) + } + if jc <= 0 { + return ipcErrorf(ipc.IpcErrorInvalid, "jc must be a positive value") + } + device.log.Verbosef("UAPI: Updating junk count") + device.junk.count = jc + + case "jmin": + jmin, err := strconv.Atoi(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jmin: %w", err) + } + if jmin <= 0 { + return ipcErrorf(ipc.IpcErrorInvalid, "jmin must be a positive value") + } + device.log.Verbosef("UAPI: Updating junk min") + device.junk.min = jmin + + case "jmax": + jmax, err := strconv.Atoi(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jmax: %w", err) + } + if jmax <= 0 { + return ipcErrorf(ipc.IpcErrorInvalid, "jmax must be a positive value") + } + device.log.Verbosef("UAPI: Updating junk max") + device.junk.max = jmax + + case "s1": + padding, err := strconv.Atoi(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s1: %w", err) + } + if padding < 0 { + return ipcErrorf(ipc.IpcErrorInvalid, "s1 must be non-negative") + } + device.log.Verbosef("UAPI: Updating s1 padding") + device.paddings.init = padding + + case "s2": + padding, err := strconv.Atoi(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s2: %w", err) + } + if padding < 0 { + return ipcErrorf(ipc.IpcErrorInvalid, "s2 must be non-negative") + } + device.log.Verbosef("UAPI: Updating s2 padding") + device.paddings.response = padding + + case "s3": + padding, err := strconv.Atoi(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s3: %w", err) + } + if padding < 0 { + return ipcErrorf(ipc.IpcErrorInvalid, "s3 must be non-negative") + } + device.log.Verbosef("UAPI: Updating s3 padding") + device.paddings.cookie = padding + + case "s4": + padding, err := strconv.Atoi(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s4: %w", err) + } + if padding < 0 { + return ipcErrorf(ipc.IpcErrorInvalid, "s4 must be non-negative") + } + device.log.Verbosef("UAPI: Updating s4 padding") + device.paddings.transport = padding + + case "h1": + header, err := newMagicHeader(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H1: %w", err) + } + device.headers.init = header + + case "h2": + header, err := newMagicHeader(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H2: %w", err) + } + device.headers.response = header + + case "h3": + header, err := newMagicHeader(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H3: %w", err) + } + device.headers.cookie = header + + case "h4": + header, err := newMagicHeader(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H4: %w", err) + } + device.headers.transport = header + + case "i1": + chain, err := newObfChain(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I1: %w", err) + } + device.ipackets[0] = chain + + case "i2": + chain, err := newObfChain(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I2: %w", err) + } + device.ipackets[1] = chain + + case "i3": + chain, err := newObfChain(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I3: %w", err) + } + device.ipackets[2] = chain + + case "i4": + chain, err := newObfChain(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I4: %w", err) + } + device.ipackets[3] = chain + + case "i5": + chain, err := newObfChain(value) + if err != nil { + return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I5: %w", err) + } + device.ipackets[4] = chain + default: return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key) } @@ -271,7 +478,10 @@ func (peer *ipcSetPeer) handlePostConfig() { } } -func (device *Device) handlePublicKeyLine(peer *ipcSetPeer, value string) error { +func (device *Device) handlePublicKeyLine( + peer *ipcSetPeer, + value string, +) error { // Load/create the peer we are configuring. var publicKey NoisePublicKey err := publicKey.FromHex(value) @@ -301,12 +511,19 @@ func (device *Device) handlePublicKeyLine(peer *ipcSetPeer, value string) error return nil } -func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error { +func (device *Device) handlePeerLine( + peer *ipcSetPeer, + key, value string, +) error { switch key { case "update_only": // allow disabling of creation if value != "true" { - return ipcErrorf(ipc.IpcErrorInvalid, "failed to set update only, invalid value: %v", value) + return ipcErrorf( + ipc.IpcErrorInvalid, + "failed to set update only, invalid value: %v", + value, + ) } if peer.created && !peer.dummy { device.RemovePeer(peer.handshake.remoteStatic) @@ -352,7 +569,11 @@ func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error secs, err := strconv.ParseUint(value, 10, 16) if err != nil { - return ipcErrorf(ipc.IpcErrorInvalid, "failed to set persistent keepalive interval: %w", err) + return ipcErrorf( + ipc.IpcErrorInvalid, + "failed to set persistent keepalive interval: %w", + err, + ) } old := peer.persistentKeepaliveInterval.Swap(uint32(secs)) @@ -363,7 +584,11 @@ func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error case "replace_allowed_ips": device.log.Verbosef("%v - UAPI: Removing all allowedips", peer.Peer) if value != "true" { - return ipcErrorf(ipc.IpcErrorInvalid, "failed to replace allowedips, invalid value: %v", value) + return ipcErrorf( + ipc.IpcErrorInvalid, + "failed to replace allowedips, invalid value: %v", + value, + ) } if peer.dummy { return nil @@ -442,7 +667,11 @@ func (device *Device) IpcHandle(socket net.Conn) { return } if nextByte != '\n' { - err = ipcErrorf(ipc.IpcErrorInvalid, "trailing character in UAPI get: %q", nextByte) + err = ipcErrorf( + ipc.IpcErrorInvalid, + "trailing character in UAPI get: %q", + nextByte, + ) break } err = device.IpcGetOperation(buffered.Writer) @@ -466,3 +695,49 @@ func (device *Device) IpcHandle(socket net.Conn) { buffered.Flush() } } + +type ipcSetDevice struct { + headers struct { + init *magicHeader + response *magicHeader + cookie *magicHeader + transport *magicHeader + } +} + +func (d *ipcSetDevice) mergeWithDevice(device *Device) error { + if d.headers.init == nil { + d.headers.init = device.headers.init + } + + if d.headers.response == nil { + d.headers.response = device.headers.response + } + + if d.headers.cookie == nil { + d.headers.cookie = device.headers.cookie + } + + if d.headers.transport == nil { + d.headers.transport = device.headers.transport + } + + headers := []*magicHeader{d.headers.init, d.headers.response, d.headers.cookie, d.headers.transport} + for i := 0; i < len(headers); i++ { + for j := i + 1; j < len(headers); j++ { + left := headers[i] + right := headers[j] + + if left.start <= right.end && right.start <= left.end { + return errors.New("headers must not overlap") + } + } + } + + device.headers.init = d.headers.init + device.headers.response = d.headers.response + device.headers.cookie = d.headers.cookie + device.headers.transport = d.headers.transport + + return nil +}