Skip to content

Commit 5aa4400

Browse files
committed
wg: do not use tun
wg is already a l2 device, it doesn't need to use tcpip
1 parent 14507a1 commit 5aa4400

2 files changed

Lines changed: 204 additions & 163 deletions

File tree

tun.go

Lines changed: 56 additions & 85 deletions
Original file line numberDiff line numberDiff line change
@@ -9,12 +9,12 @@ import (
99
"net/netip"
1010
"os"
1111
// "regexp"
12+
mathrand "math/rand"
1213
"strconv"
1314
"sync"
1415
"sync/atomic"
1516
"syscall"
1617
"time"
17-
mathrand "math/rand"
1818

1919
// "github.com/google/gopacket"
2020
// "github.com/google/gopacket/layers"
@@ -32,71 +32,60 @@ import (
3232
"gvisor.dev/gvisor/pkg/tcpip/transport/udp"
3333
// "gvisor.dev/gvisor/pkg/waiter"
3434

35-
36-
"github.com/urnetwork/glog"
3735
"github.com/urnetwork/connect"
36+
"github.com/urnetwork/glog"
3837
)
3938

40-
4139
// const DefaultChannelSize = 64
4240
// const DefaultProxySequenceSize = 64
4341
// const DefaultWriteTimeout = 15 * time.Second
4442

4543
func DefaultTunSettings() *TunSettings {
4644
return &TunSettings{
47-
ChannelSize: 64,
45+
ChannelSize: 64,
4846
ProxySequenceSize: 64,
49-
Mtu: 1440,
47+
Mtu: 1440,
5048
}
5149
}
5250

53-
54-
55-
5651
type TunSettings struct {
57-
ChannelSize int
52+
ChannelSize int
5853
ProxySequenceSize int
59-
Mtu int
54+
Mtu int
6055
}
6156

62-
63-
64-
var tunStack = sync.OnceValue(func()(*stack.Stack) {
57+
var tunStack = sync.OnceValue(func() *stack.Stack {
6558
opts := stack.Options{
66-
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocolWithOptions(ipv4.Options{AllowExternalLoopbackTraffic:true})},
59+
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocolWithOptions(ipv4.Options{AllowExternalLoopbackTraffic: true})},
6760
TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol, icmp.NewProtocol4},
6861
HandleLocal: true,
6962
}
7063
return stack.New(opts)
7164

7265
})
7366

74-
7567
// FIXME this should be a pool where closed nic ids and addrs can be reused
7668
var nicIdCounter atomic.Uint32
77-
var localIpv4AddressGenerator = sync.OnceValue(func()(*connect.AddrGenerator) {
69+
var localIpv4AddressGenerator = sync.OnceValue(func() *connect.AddrGenerator {
7870
prefix := netip.MustParsePrefix("169.254.0.0/16")
7971
return connect.NewAddrGenerator(prefix)
8072
})
8173

82-
8374
type Tun struct {
84-
ctx context.Context
75+
ctx context.Context
8576
cancel context.CancelFunc
8677

8778
settings *TunSettings
8879

89-
ep *channel.Endpoint
90-
stack *stack.Stack
91-
nicId tcpip.NICID
92-
receivePacket chan []byte
80+
ep *channel.Endpoint
81+
stack *stack.Stack
82+
nicId tcpip.NICID
83+
receivePacket chan []byte
9384
// mtu int
9485
// registeredAddresses map[netip.Addr]bool
95-
dohResolver *connect.DohCache
86+
dohResolver *connect.DohCache
9687

9788
stateLock sync.Mutex
98-
customReceivePacket chan []byte
99-
10089
}
10190

10291
func CreateTunWithDefaults(ctx context.Context) (*Tun, error) {
@@ -122,15 +111,15 @@ func CreateTunWithResolver(ctx context.Context, settings *TunSettings, dnsResolv
122111
}
123112

124113
ep := channel.New(settings.ChannelSize, uint32(settings.Mtu), tcpip.LinkAddress(fmt.Sprintf("%x", nicId)))
125-
114+
126115
tun := &Tun{
127-
ctx: cancelCtx,
128-
cancel: cancel,
129-
settings: settings,
130-
ep: ep,
131-
stack: tunStack(),
132-
nicId: nicId,
133-
receivePacket: make(chan []byte, settings.ProxySequenceSize),
116+
ctx: cancelCtx,
117+
cancel: cancel,
118+
settings: settings,
119+
ep: ep,
120+
stack: tunStack(),
121+
nicId: nicId,
122+
receivePacket: make(chan []byte, settings.ProxySequenceSize),
134123
}
135124

136125
dohSettings := connect.DefaultDohSettings()
@@ -160,7 +149,7 @@ func CreateTunWithResolver(ctx context.Context, settings *TunSettings, dnsResolv
160149
Protocol: protoNumber,
161150
AddressWithPrefix: tcpip.AddrFromSlice(ip.AsSlice()).WithPrefix(),
162151
}
163-
152+
164153
if tcpipErr := tun.stack.AddProtocolAddress(nicId, protoAddr, stack.AddressProperties{}); tcpipErr != nil {
165154
return nil, fmt.Errorf("Could not create add nic address err=%s", tcpipErr)
166155
}
@@ -172,31 +161,13 @@ func CreateTunWithResolver(ctx context.Context, settings *TunSettings, dnsResolv
172161
return tun, nil
173162
}
174163

175-
176-
func (self *Tun) SetReceive(customReceivePacket chan []byte) {
177-
self.stateLock.Lock()
178-
defer self.stateLock.Unlock()
179-
180-
self.customReceivePacket = customReceivePacket
181-
}
182-
183-
func (self *Tun) receive() chan []byte {
184-
self.stateLock.Lock()
185-
defer self.stateLock.Unlock()
186-
187-
if self.customReceivePacket != nil {
188-
return self.customReceivePacket
189-
}
190-
return self.receivePacket
191-
}
192-
193164
func (self *Tun) DohCache() *connect.DohCache {
194165
return self.dohResolver
195166
}
196167

197168
func (self *Tun) Read() ([]byte, error) {
198169
select {
199-
case <- self.ctx.Done():
170+
case <-self.ctx.Done():
200171
return nil, fmt.Errorf("Done")
201172
case m, ok := <-self.receivePacket:
202173
if !ok {
@@ -243,12 +214,12 @@ func (self *Tun) WriteNotify() {
243214
pkt.DecRef()
244215

245216
select {
246-
case <- self.ctx.Done():
217+
case <-self.ctx.Done():
247218
connect.MessagePoolReturn(packet)
248-
case self.receive() <- packet:
249-
// case <-time.After(DefaultWriteTimeout):
250-
// // drop
251-
// connect.MessagePoolReturn(packet)
219+
case self.receivePacket <- packet:
220+
// case <-time.After(DefaultWriteTimeout):
221+
// // drop
222+
// connect.MessagePoolReturn(packet)
252223
}
253224
}
254225

@@ -274,8 +245,8 @@ func (self *Tun) dialCtx(ctx context.Context) context.Context {
274245
go func() {
275246
defer dialCancel()
276247
select {
277-
case <- ctx.Done():
278-
case <- self.ctx.Done():
248+
case <-ctx.Done():
249+
case <-self.ctx.Done():
279250
}
280251
}()
281252
return dialCtx
@@ -314,7 +285,7 @@ func (self *Tun) DialContext(ctx context.Context, network string, address string
314285
if err != nil {
315286
return nil, err
316287
}
317-
288+
318289
var addrs []netip.Addr
319290
if addr, err := netip.ParseAddr(host); err == nil {
320291
// address is ip:port
@@ -337,30 +308,30 @@ func (self *Tun) DialContext(ctx context.Context, network string, address string
337308

338309
// var returnErr error
339310
// for _, addr := range addrs {
340-
addrPort := netip.AddrPortFrom(addr, uint16(port))
341-
342-
switch network {
343-
case "tcp", "tcp4", "tcp6":
344-
fa, pn := self.convertToFullAddr(addrPort)
345-
conn, err := gonet.DialTCP(self.stack, fa, pn)
346-
if err == nil {
347-
glog.V(1).Infof("[tun]tcp connect (%s)->%s success\n", host, addrPort)
348-
return conn, nil
349-
}
350-
glog.V(1).Infof("[tun]tcp connect (%s)->%s err = %s\n", host, addrPort, err)
351-
return nil, err
352-
case "udp", "udp4", "udp6":
353-
fa, pn := self.convertToFullAddr(addrPort)
354-
conn, err := gonet.DialUDP(self.stack, nil, &fa, pn)
355-
if err == nil {
356-
glog.V(1).Infof("[tun]udp connect (%s)->%s success\n", host, addrPort)
357-
return conn, nil
358-
}
359-
glog.V(1).Infof("[tun]tcp connect (%s)->%s err = %s\n", host, addrPort, err)
360-
return nil, err
361-
default:
362-
return nil, fmt.Errorf("Unsupported network %s", network)
311+
addrPort := netip.AddrPortFrom(addr, uint16(port))
312+
313+
switch network {
314+
case "tcp", "tcp4", "tcp6":
315+
fa, pn := self.convertToFullAddr(addrPort)
316+
conn, err := gonet.DialTCP(self.stack, fa, pn)
317+
if err == nil {
318+
glog.V(1).Infof("[tun]tcp connect (%s)->%s success\n", host, addrPort)
319+
return conn, nil
320+
}
321+
glog.V(1).Infof("[tun]tcp connect (%s)->%s err = %s\n", host, addrPort, err)
322+
return nil, err
323+
case "udp", "udp4", "udp6":
324+
fa, pn := self.convertToFullAddr(addrPort)
325+
conn, err := gonet.DialUDP(self.stack, nil, &fa, pn)
326+
if err == nil {
327+
glog.V(1).Infof("[tun]udp connect (%s)->%s success\n", host, addrPort)
328+
return conn, nil
363329
}
330+
glog.V(1).Infof("[tun]tcp connect (%s)->%s err = %s\n", host, addrPort, err)
331+
return nil, err
332+
default:
333+
return nil, fmt.Errorf("Unsupported network %s", network)
334+
}
364335
// }
365336

366337
// return nil, returnErr

0 commit comments

Comments
 (0)