Remove unused
This commit is contained in:
parent
e620c55272
commit
f853bfc5c5
34 changed files with 0 additions and 6510 deletions
77
README.md
77
README.md
|
|
@ -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.
|
|
||||||
|
|
@ -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])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -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
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
@ -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")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
|
|
@ -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
|
|
||||||
}
|
|
||||||
|
|
@ -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,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -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
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}()
|
|
||||||
}
|
|
||||||
|
|
@ -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()
|
|
||||||
}
|
|
||||||
|
|
@ -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
|
|
||||||
|
|
@ -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
|
|
||||||
|
|
@ -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()
|
|
||||||
}
|
|
||||||
6
go.mod
6
go.mod
|
|
@ -7,10 +7,4 @@ require (
|
||||||
golang.org/x/net v0.39.0
|
golang.org/x/net v0.39.0
|
||||||
golang.org/x/sys v0.32.0
|
golang.org/x/sys v0.32.0
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
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
|
|
||||||
)
|
)
|
||||||
|
|
|
||||||
6
go.sum
6
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 h1:kJNSjF/Xp7kU0iB2Z+9viTPMW4EqqsrywMXLJOOsXSE=
|
||||||
golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc=
|
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 h1:ZCu7HMWDxpXpaiKdhzIfaltL9Lp31x/3fCP11bc6/fY=
|
||||||
golang.org/x/net v0.39.0/go.mod h1:X7NRbYVEA+ewNkCNyJ513WmMdQ3BineSwVtN2zD/d+E=
|
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 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20=
|
||||||
golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
|
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 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
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=
|
|
||||||
|
|
|
||||||
|
|
@ -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()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
268
main.go
268
main.go
|
|
@ -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 <Jason@zx2c4.com>.\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")
|
|
||||||
}
|
|
||||||
|
|
@ -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 <https://git.zx2c4.com/wireguard-windows/>, 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")
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
425
tests/netns.sh
425
tests/netns.sh
|
|
@ -1,425 +0,0 @@
|
||||||
#!/bin/bash
|
|
||||||
|
|
||||||
# Copyright (C) 2015-2017 Jason A. Donenfeld <Jason@zx2c4.com>. 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 <path to wireguard-go>
|
|
||||||
|
|
||||||
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."
|
|
||||||
|
|
@ -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))
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -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))
|
|
||||||
}
|
|
||||||
|
|
@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -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))
|
|
||||||
}
|
|
||||||
1057
tun/netstack/tun.go
1057
tun/netstack/tun.go
File diff suppressed because it is too large
Load diff
|
|
@ -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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -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
|
|
||||||
}
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue