| 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 | } |