main
go 269 lines 6.59 KB
Raw
1 //go:build linux || darwin || windows || freebsd || openbsd
2
3 package overlay
4
5 import (
6 "context"
7 "errors"
8 "fmt"
9 "net"
10 "net/netip"
11 "strconv"
12 "strings"
13 "sync"
14 "time"
15
16 "golang.zx2c4.com/wireguard/conn"
17 "golang.zx2c4.com/wireguard/device"
18 "golang.zx2c4.com/wireguard/tun/netstack"
19
20 "github.com/gosuda/portal-tunnel/v2/portal/identity"
21 "github.com/gosuda/portal-tunnel/v2/types"
22 )
23
24 const defaultEndpointResolveTTL = 3 * time.Second
25
26 type stack struct {
27 device *device.Device
28 net *netstack.Net
29 overlayIP netip.Addr
30
31 applyMu sync.Mutex
32 mu sync.Mutex
33 closed bool
34 peerEndpoints map[string]string
35 peerConfig string
36 }
37
38 func newStack(cfg Config) (*stack, error) {
39 canonicalPrivateKey, err := identity.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
40 if err != nil {
41 return nil, fmt.Errorf("normalize wireguard private key: %w", err)
42 }
43
44 listenPort := cfg.ListenPort
45 if listenPort <= 0 || listenPort > 65535 {
46 return nil, errors.New("wireguard listen port is invalid")
47 }
48
49 overlayIPv4, err := identity.DeriveWireGuardOverlayIPv4(cfg.PublicKey)
50 if err != nil {
51 return nil, fmt.Errorf("derive overlay ipv4: %w", err)
52 }
53 overlayIP, err := netip.ParseAddr(overlayIPv4)
54 if err != nil || !overlayIP.Is4() {
55 return nil, errors.New("overlay ipv4 must be a valid IPv4 address")
56 }
57
58 tunDevice, network, err := netstack.CreateNetTUN([]netip.Addr{overlayIP}, nil, DefaultMTU)
59 if err != nil {
60 return nil, fmt.Errorf("create netstack tun: %w", err)
61 }
62
63 wgDevice := device.NewDevice(tunDevice, conn.NewDefaultBind(), device.NewLogger(device.LogLevelError, "portal-wg"))
64 privateKeyHex, err := identity.WireGuardKeyHex(canonicalPrivateKey)
65 if err != nil {
66 wgDevice.Close()
67 <-wgDevice.Wait()
68 return nil, err
69 }
70
71 config := fmt.Sprintf("private_key=%s\nlisten_port=%d\n", privateKeyHex, listenPort)
72 if err := wgDevice.IpcSet(config); err != nil {
73 wgDevice.Close()
74 <-wgDevice.Wait()
75 return nil, fmt.Errorf("configure wireguard device: %w", err)
76 }
77 if err := wgDevice.Up(); err != nil {
78 wgDevice.Close()
79 <-wgDevice.Wait()
80 return nil, fmt.Errorf("bring wireguard device up: %w", err)
81 }
82
83 return &stack{
84 device: wgDevice,
85 net: network,
86 overlayIP: overlayIP,
87 peerEndpoints: map[string]string{},
88 }, nil
89 }
90
91 func (s *stack) ListenTCP(port int) (net.Listener, error) {
92 if s == nil || s.net == nil {
93 return nil, errors.New("wireguard is not initialized")
94 }
95 return s.net.ListenTCP(&net.TCPAddr{
96 IP: net.ParseIP(s.overlayIP.String()),
97 Port: port,
98 })
99 }
100
101 func (s *stack) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
102 if s == nil || s.net == nil {
103 return nil, errors.New("wireguard is not initialized")
104 }
105 switch network {
106 case "tcp", "tcp4", "tcp6":
107 default:
108 return nil, fmt.Errorf("unsupported network %q", network)
109 }
110
111 host, portText, err := net.SplitHostPort(address)
112 if err != nil {
113 return nil, err
114 }
115 ip, err := netip.ParseAddr(strings.Trim(host, "[]"))
116 if err != nil {
117 return nil, err
118 }
119 port, err := strconv.Atoi(portText)
120 if err != nil || port <= 0 || port > 65535 {
121 return nil, errors.New("invalid tcp port")
122 }
123 return s.net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(port)))
124 }
125
126 func (s *stack) ApplyPeers(peers []types.RelayDescriptor) error {
127 if s == nil || s.device == nil {
128 return errors.New("wireguard is not initialized")
129 }
130 s.applyMu.Lock()
131 defer s.applyMu.Unlock()
132 s.mu.Lock()
133 if s.closed {
134 s.mu.Unlock()
135 return net.ErrClosed
136 }
137 s.mu.Unlock()
138
139 var builder strings.Builder
140 builder.WriteString("replace_peers=true\n")
141 var warnErr error
142 nextPeerEndpoints := map[string]string{}
143
144 for _, peer := range peers {
145 peerKey := strings.TrimSpace(peer.WireGuardPublicKey)
146 overlayIPv4, err := identity.DeriveWireGuardOverlayIPv4(peer.WireGuardPublicKey)
147 if err != nil {
148 continue
149 }
150 wireGuardEndpoint, err := identity.RelayWireGuardEndpoint(peer)
151 if err != nil {
152 continue
153 }
154 publicKeyHex, err := identity.WireGuardKeyHex(peer.WireGuardPublicKey)
155 if err != nil {
156 return fmt.Errorf("normalize peer %q public key: %w", peerKey, err)
157 }
158
159 resolvedEndpoint, err := resolvePeerEndpoint(wireGuardEndpoint)
160 if err != nil {
161 s.mu.Lock()
162 currentEndpoint := s.peerEndpoints[publicKeyHex]
163 s.mu.Unlock()
164 if currentEndpoint != "" {
165 warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w; using current endpoint %q", peerKey, err, currentEndpoint))
166 resolvedEndpoint = currentEndpoint
167 } else {
168 warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w", peerKey, err))
169 continue
170 }
171 }
172
173 builder.WriteString("public_key=")
174 builder.WriteString(publicKeyHex)
175 builder.WriteByte('\n')
176 builder.WriteString("endpoint=")
177 builder.WriteString(resolvedEndpoint)
178 builder.WriteByte('\n')
179 nextPeerEndpoints[publicKeyHex] = resolvedEndpoint
180 builder.WriteString("allowed_ip=")
181 builder.WriteString(overlayIPv4)
182 builder.WriteString("/32\n")
183 if DefaultPersistentKeepalive > 0 {
184 builder.WriteString("persistent_keepalive_interval=")
185 builder.WriteString(strconv.Itoa(DefaultPersistentKeepalive))
186 builder.WriteByte('\n')
187 }
188 }
189
190 config := builder.String()
191 s.mu.Lock()
192 if s.peerConfig == config {
193 s.mu.Unlock()
194 return warnErr
195 }
196 s.mu.Unlock()
197
198 if err := s.device.IpcSet(config); err != nil {
199 return err
200 }
201 s.mu.Lock()
202 s.peerEndpoints = nextPeerEndpoints
203 s.peerConfig = config
204 s.mu.Unlock()
205 return warnErr
206 }
207
208 func resolvePeerEndpoint(raw string) (string, error) {
209 endpoint := strings.TrimSpace(raw)
210 if endpoint == "" {
211 return "", errors.New("wireguard endpoint is required")
212 }
213
214 host, port, err := net.SplitHostPort(endpoint)
215 if err != nil {
216 return "", err
217 }
218
219 host = strings.Trim(host, "[]")
220 if host == "" {
221 return "", errors.New("wireguard endpoint host is required")
222 }
223
224 if ip, err := netip.ParseAddr(host); err == nil {
225 return net.JoinHostPort(ip.String(), port), nil
226 }
227
228 ctx, cancel := context.WithTimeout(context.Background(), defaultEndpointResolveTTL)
229 defer cancel()
230
231 addrs, err := net.DefaultResolver.LookupNetIP(ctx, "ip", host)
232 if err != nil {
233 return "", fmt.Errorf("lookup %q: %w", host, err)
234 }
235 if len(addrs) == 0 {
236 return "", fmt.Errorf("lookup %q: no IP addresses found", host)
237 }
238
239 selected := addrs[0]
240 for _, addr := range addrs {
241 if addr.Is4() {
242 selected = addr
243 break
244 }
245 }
246 return net.JoinHostPort(selected.String(), port), nil
247 }
248
249 func (s *stack) Close() error {
250 if s == nil || s.device == nil {
251 return nil
252 }
253
254 s.applyMu.Lock()
255 defer s.applyMu.Unlock()
256
257 s.mu.Lock()
258 if s.closed {
259 s.mu.Unlock()
260 return nil
261 }
262 s.closed = true
263 device := s.device
264 s.mu.Unlock()
265
266 device.Close()
267 <-device.Wait()
268 return nil
269 }