mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
chore: switch to new wgtunnel via tunnelsdk (#6489)
This commit is contained in:
+37
-13
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/go-ping/ping"
|
||||
"golang.org/x/exp/slices"
|
||||
"golang.org/x/sync/errgroup"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/cryptorand"
|
||||
)
|
||||
@@ -19,13 +20,11 @@ type Region struct {
|
||||
}
|
||||
|
||||
type Node struct {
|
||||
ID int `json:"id"`
|
||||
RegionID int `json:"region_id"`
|
||||
HostnameHTTPS string `json:"hostname_https"`
|
||||
HostnameWireguard string `json:"hostname_wireguard"`
|
||||
WireguardPort uint16 `json:"wireguard_port"`
|
||||
ID int `json:"id"`
|
||||
RegionID int `json:"region_id"`
|
||||
HostnameHTTPS string `json:"hostname_https"`
|
||||
|
||||
AvgLatency time.Duration `json:"avg_latency"`
|
||||
AvgLatency time.Duration `json:"-"`
|
||||
}
|
||||
|
||||
var Regions = []Region{
|
||||
@@ -34,28 +33,53 @@ var Regions = []Region{
|
||||
LocationName: "US East Pittsburgh",
|
||||
Nodes: []Node{
|
||||
{
|
||||
ID: 1,
|
||||
RegionID: 0,
|
||||
HostnameHTTPS: "pit-1.try.coder.app",
|
||||
HostnameWireguard: "pit-1.try.coder.app",
|
||||
WireguardPort: 55551,
|
||||
ID: 1,
|
||||
RegionID: 0,
|
||||
HostnameHTTPS: "pit-1.try.coder.app",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
func FindClosestNode() (Node, error) {
|
||||
// Nodes returns a list of nodes to use for the tunnel. It will pick a random
|
||||
// node from each region.
|
||||
//
|
||||
// If a customNode is provided, it will be returned as the only node with ID
|
||||
// 9999.
|
||||
func Nodes(customTunnelHost string) ([]Node, error) {
|
||||
nodes := []Node{}
|
||||
|
||||
if customTunnelHost != "" {
|
||||
return []Node{
|
||||
{
|
||||
ID: 9999,
|
||||
RegionID: 9999,
|
||||
HostnameHTTPS: customTunnelHost,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
for _, region := range Regions {
|
||||
// Pick a random node from each region.
|
||||
i, err := cryptorand.Intn(len(region.Nodes))
|
||||
if err != nil {
|
||||
return Node{}, err
|
||||
return []Node{}, err
|
||||
}
|
||||
nodes = append(nodes, region.Nodes[i])
|
||||
}
|
||||
|
||||
return nodes, nil
|
||||
}
|
||||
|
||||
// FindClosestNode pings each node and returns the one with the lowest latency.
|
||||
func FindClosestNode(nodes []Node) (Node, error) {
|
||||
if len(nodes) == 0 {
|
||||
return Node{}, xerrors.New("no wgtunnel nodes")
|
||||
}
|
||||
|
||||
// Copy the nodes so we don't mutate the original.
|
||||
nodes = append([]Node{}, nodes...)
|
||||
|
||||
var (
|
||||
nodesMu sync.Mutex
|
||||
eg = errgroup.Group{}
|
||||
|
||||
+49
-181
@@ -1,134 +1,52 @@
|
||||
package devtunnel
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/briandowns/spinner"
|
||||
"golang.org/x/xerrors"
|
||||
"golang.zx2c4.com/wireguard/conn"
|
||||
"golang.zx2c4.com/wireguard/device"
|
||||
"golang.zx2c4.com/wireguard/tun/netstack"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/cli/cliui"
|
||||
"github.com/coder/coder/cryptorand"
|
||||
"github.com/coder/wgtunnel/tunnelsdk"
|
||||
)
|
||||
|
||||
type Tunnel struct {
|
||||
URL string
|
||||
Listener net.Listener
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
Version int `json:"version"`
|
||||
PrivateKey device.NoisePrivateKey `json:"private_key"`
|
||||
PublicKey device.NoisePublicKey `json:"public_key"`
|
||||
Version tunnelsdk.TunnelVersion `json:"version"`
|
||||
PrivateKey device.NoisePrivateKey `json:"private_key"`
|
||||
PublicKey device.NoisePublicKey `json:"public_key"`
|
||||
|
||||
Tunnel Node `json:"tunnel"`
|
||||
|
||||
// Used in testing. Normally this is nil, indicating to use DefaultClient.
|
||||
HTTPClient *http.Client `json:"-"`
|
||||
}
|
||||
type configExt struct {
|
||||
Version int `json:"-"`
|
||||
PrivateKey device.NoisePrivateKey `json:"-"`
|
||||
PublicKey device.NoisePublicKey `json:"public_key"`
|
||||
|
||||
Tunnel Node `json:"-"`
|
||||
|
||||
// Used in testing. Normally this is nil, indicating to use DefaultClient.
|
||||
HTTPClient *http.Client `json:"-"`
|
||||
}
|
||||
|
||||
// NewWithConfig calls New with the given config. For documentation, see New.
|
||||
func NewWithConfig(ctx context.Context, logger slog.Logger, cfg Config) (*Tunnel, <-chan error, error) {
|
||||
server, routineEnd, err := startUpdateRoutine(ctx, logger, cfg)
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("start update routine: %w", err)
|
||||
func NewWithConfig(ctx context.Context, logger slog.Logger, cfg Config) (*tunnelsdk.Tunnel, error) {
|
||||
u := &url.URL{
|
||||
Scheme: "https",
|
||||
Host: cfg.Tunnel.HostnameHTTPS,
|
||||
}
|
||||
|
||||
tun, tnet, err := netstack.CreateNetTUN(
|
||||
[]netip.Addr{server.ClientIP},
|
||||
[]netip.Addr{netip.AddrFrom4([4]byte{1, 1, 1, 1})},
|
||||
1280,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("create net TUN: %w", err)
|
||||
c := tunnelsdk.New(u)
|
||||
if cfg.HTTPClient != nil {
|
||||
c.HTTPClient = cfg.HTTPClient
|
||||
}
|
||||
|
||||
wgip, err := net.ResolveIPAddr("ip", cfg.Tunnel.HostnameWireguard)
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("resolve endpoint: %w", err)
|
||||
}
|
||||
// In IPv6, we need to enclose the address to in [] before passing to wireguard's endpoint key, like
|
||||
// [2001:abcd::1]:8888. We'll use netip.AddrPort to correctly handle this.
|
||||
wgAddr, err := netip.ParseAddr(wgip.String())
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("parse address: %w", err)
|
||||
}
|
||||
wgEndpoint := netip.AddrPortFrom(wgAddr, cfg.Tunnel.WireguardPort)
|
||||
|
||||
dlog := &device.Logger{
|
||||
Verbosef: slog.Stdlib(ctx, logger, slog.LevelDebug).Printf,
|
||||
Errorf: slog.Stdlib(ctx, logger, slog.LevelError).Printf,
|
||||
}
|
||||
dev := device.NewDevice(tun, conn.NewDefaultBind(), dlog)
|
||||
err = dev.IpcSet(fmt.Sprintf(`private_key=%s
|
||||
public_key=%s
|
||||
endpoint=%s
|
||||
persistent_keepalive_interval=21
|
||||
allowed_ip=%s/128`,
|
||||
hex.EncodeToString(cfg.PrivateKey[:]),
|
||||
server.ServerPublicKey,
|
||||
wgEndpoint.String(),
|
||||
server.ServerIP.String(),
|
||||
))
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("configure wireguard ipc: %w", err)
|
||||
}
|
||||
|
||||
err = dev.Up()
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("wireguard device up: %w", err)
|
||||
}
|
||||
|
||||
wgListen, err := tnet.ListenTCP(&net.TCPAddr{Port: 8090})
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("wireguard device listen: %w", err)
|
||||
}
|
||||
|
||||
ch := make(chan error, 1)
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
_ = wgListen.Close()
|
||||
// We need to remove peers before closing to avoid a race condition between dev.Close() and the peer
|
||||
// goroutines which results in segfault.
|
||||
dev.RemoveAllPeers()
|
||||
dev.Close()
|
||||
<-routineEnd
|
||||
close(ch)
|
||||
|
||||
case <-dev.Wait():
|
||||
close(ch)
|
||||
}
|
||||
}()
|
||||
|
||||
return &Tunnel{
|
||||
URL: fmt.Sprintf("https://%s", server.Hostname),
|
||||
Listener: wgListen,
|
||||
}, ch, nil
|
||||
return c.LaunchTunnel(ctx, tunnelsdk.TunnelConfig{
|
||||
Log: logger,
|
||||
Version: cfg.Version,
|
||||
PrivateKey: tunnelsdk.FromNoisePrivateKey(cfg.PrivateKey),
|
||||
})
|
||||
}
|
||||
|
||||
// New creates a tunnel with a public URL and returns a listener for incoming
|
||||
@@ -136,82 +54,18 @@ allowed_ip=%s/128`,
|
||||
// Tunnel configuration is cached in the user's config directory. Successive
|
||||
// calls to New will always use the same URL. If multiple public URLs in
|
||||
// parallel are required, use NewWithConfig.
|
||||
func New(ctx context.Context, logger slog.Logger) (*Tunnel, <-chan error, error) {
|
||||
cfg, err := readOrGenerateConfig()
|
||||
//
|
||||
// This uses https://github.com/coder/wgtunnel as the server and client
|
||||
// implementation.
|
||||
func New(ctx context.Context, logger slog.Logger, customTunnelHost string) (*tunnelsdk.Tunnel, error) {
|
||||
cfg, err := readOrGenerateConfig(customTunnelHost)
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("read or generate config: %w", err)
|
||||
return nil, xerrors.Errorf("read or generate config: %w", err)
|
||||
}
|
||||
|
||||
return NewWithConfig(ctx, logger, cfg)
|
||||
}
|
||||
|
||||
func startUpdateRoutine(ctx context.Context, logger slog.Logger, cfg Config) (ServerResponse, <-chan struct{}, error) {
|
||||
// Ensure we send the first config before spawning in the background.
|
||||
res, err := sendConfigToServer(ctx, cfg)
|
||||
if err != nil {
|
||||
return ServerResponse{}, nil, xerrors.Errorf("send config to server: %w", err)
|
||||
}
|
||||
|
||||
endCh := make(chan struct{})
|
||||
go func() {
|
||||
defer close(endCh)
|
||||
ticker := time.NewTicker(10 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
|
||||
case <-ticker.C:
|
||||
}
|
||||
|
||||
_, err := sendConfigToServer(ctx, cfg)
|
||||
if err != nil {
|
||||
logger.Debug(ctx, "send tunnel config to server", slog.Error(err))
|
||||
}
|
||||
}
|
||||
}()
|
||||
return res, endCh, nil
|
||||
}
|
||||
|
||||
type ServerResponse struct {
|
||||
Hostname string `json:"hostname"`
|
||||
ServerIP netip.Addr `json:"server_ip"`
|
||||
ServerPublicKey string `json:"server_public_key"` // hex
|
||||
ClientIP netip.Addr `json:"client_ip"`
|
||||
}
|
||||
|
||||
func sendConfigToServer(ctx context.Context, cfg Config) (ServerResponse, error) {
|
||||
raw, err := json.Marshal(configExt(cfg))
|
||||
if err != nil {
|
||||
return ServerResponse{}, xerrors.Errorf("marshal config: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", "https://"+cfg.Tunnel.HostnameHTTPS+"/tun", bytes.NewReader(raw))
|
||||
if err != nil {
|
||||
return ServerResponse{}, xerrors.Errorf("new request: %w", err)
|
||||
}
|
||||
|
||||
client := http.DefaultClient
|
||||
if cfg.HTTPClient != nil {
|
||||
client = cfg.HTTPClient
|
||||
}
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
return ServerResponse{}, xerrors.Errorf("do request: %w", err)
|
||||
}
|
||||
defer res.Body.Close()
|
||||
|
||||
var resp ServerResponse
|
||||
err = json.NewDecoder(res.Body).Decode(&resp)
|
||||
if err != nil {
|
||||
return ServerResponse{}, xerrors.Errorf("decode response: %w", err)
|
||||
}
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func cfgPath() (string, error) {
|
||||
cfgDir, err := os.UserConfigDir()
|
||||
if err != nil {
|
||||
@@ -227,7 +81,7 @@ func cfgPath() (string, error) {
|
||||
return filepath.Join(cfgDir, "devtunnel"), nil
|
||||
}
|
||||
|
||||
func readOrGenerateConfig() (Config, error) {
|
||||
func readOrGenerateConfig(customTunnelHost string) (Config, error) {
|
||||
cfgFi, err := cfgPath()
|
||||
if err != nil {
|
||||
return Config{}, xerrors.Errorf("get config path: %w", err)
|
||||
@@ -236,7 +90,7 @@ func readOrGenerateConfig() (Config, error) {
|
||||
fi, err := os.ReadFile(cfgFi)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
cfg, err := GenerateConfig()
|
||||
cfg, err := GenerateConfig(customTunnelHost)
|
||||
if err != nil {
|
||||
return Config{}, xerrors.Errorf("generate config: %w", err)
|
||||
}
|
||||
@@ -264,7 +118,7 @@ func readOrGenerateConfig() (Config, error) {
|
||||
_, _ = fmt.Println(cliui.Styles.Error.Render("Upgrading you to the new version now. You will need to rebuild running workspaces."))
|
||||
_, _ = fmt.Println()
|
||||
|
||||
cfg, err := GenerateConfig()
|
||||
cfg, err := GenerateConfig(customTunnelHost)
|
||||
if err != nil {
|
||||
return Config{}, xerrors.Errorf("generate config: %w", err)
|
||||
}
|
||||
@@ -280,20 +134,29 @@ func readOrGenerateConfig() (Config, error) {
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func GenerateConfig() (Config, error) {
|
||||
priv, err := wgtypes.GeneratePrivateKey()
|
||||
func GenerateConfig(customTunnelHost string) (Config, error) {
|
||||
priv, err := tunnelsdk.GeneratePrivateKey()
|
||||
if err != nil {
|
||||
return Config{}, xerrors.Errorf("generate private key: %w", err)
|
||||
}
|
||||
pub := priv.PublicKey()
|
||||
privNoisePublicKey, err := priv.NoisePrivateKey()
|
||||
if err != nil {
|
||||
return Config{}, xerrors.Errorf("generate noise private key: %w", err)
|
||||
}
|
||||
pubNoisePublicKey := priv.NoisePublicKey()
|
||||
|
||||
spin := spinner.New(spinner.CharSets[39], 350*time.Millisecond)
|
||||
spin.Suffix = " Finding the closest tunnel region..."
|
||||
spin.Start()
|
||||
|
||||
node, err := FindClosestNode()
|
||||
nodes, err := Nodes(customTunnelHost)
|
||||
if err != nil {
|
||||
// If we fail to find the closest node, default to US East.
|
||||
return Config{}, xerrors.Errorf("get nodes: %w", err)
|
||||
}
|
||||
node, err := FindClosestNode(nodes)
|
||||
if err != nil {
|
||||
// If we fail to find the closest node, default to a random node from
|
||||
// the first region.
|
||||
region := Regions[0]
|
||||
n, _ := cryptorand.Intn(len(region.Nodes))
|
||||
node = region.Nodes[n]
|
||||
@@ -302,16 +165,21 @@ func GenerateConfig() (Config, error) {
|
||||
_, _ = fmt.Println("Defaulting to", Regions[0].LocationName)
|
||||
}
|
||||
|
||||
locationName := "Unknown"
|
||||
if node.RegionID < len(Regions) {
|
||||
locationName = Regions[node.RegionID].LocationName
|
||||
}
|
||||
|
||||
spin.Stop()
|
||||
_, _ = fmt.Printf("Using tunnel in %s with latency %s.\n",
|
||||
cliui.Styles.Keyword.Render(Regions[node.RegionID].LocationName),
|
||||
cliui.Styles.Keyword.Render(locationName),
|
||||
cliui.Styles.Code.Render(node.AvgLatency.String()),
|
||||
)
|
||||
|
||||
return Config{
|
||||
Version: 1,
|
||||
PrivateKey: device.NoisePrivateKey(priv),
|
||||
PublicKey: device.NoisePublicKey(pub),
|
||||
Version: tunnelsdk.TunnelVersion2,
|
||||
PrivateKey: privNoisePublicKey,
|
||||
PublicKey: pubNoisePublicKey,
|
||||
Tunnel: node,
|
||||
}, nil
|
||||
}
|
||||
|
||||
+199
-169
@@ -2,14 +2,16 @@ package devtunnel_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"encoding/base32"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -18,26 +20,12 @@ import (
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.zx2c4.com/wireguard/conn"
|
||||
"golang.zx2c4.com/wireguard/device"
|
||||
"golang.zx2c4.com/wireguard/tun/netstack"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
"github.com/coder/coder/coderd/devtunnel"
|
||||
"github.com/coder/coder/testutil"
|
||||
)
|
||||
|
||||
const (
|
||||
ipByte1 = 0xfc
|
||||
ipByte2 = 0xca
|
||||
wgPort = 48732
|
||||
)
|
||||
|
||||
var (
|
||||
serverIP = netip.AddrFrom16([16]byte{ipByte1, ipByte2, 15: 0x1})
|
||||
dnsIP = netip.AddrFrom4([4]byte{1, 1, 1, 1})
|
||||
clientIP = netip.AddrFrom16([16]byte{ipByte1, ipByte2, 15: 0x2})
|
||||
"github.com/coder/wgtunnel/tunneld"
|
||||
"github.com/coder/wgtunnel/tunnelsdk"
|
||||
)
|
||||
|
||||
// The tunnel leaks a few goroutines that aren't impactful to production scenarios.
|
||||
@@ -45,194 +33,236 @@ var (
|
||||
// goleak.VerifyTestMain(m)
|
||||
// }
|
||||
|
||||
// TestTunnel cannot run in parallel because we hardcode the UDP port used by the wireguard server.
|
||||
// nolint: paralleltest
|
||||
func TestTunnel(t *testing.T) {
|
||||
ctx, cancelTun := context.WithCancel(context.Background())
|
||||
defer cancelTun()
|
||||
t.Parallel()
|
||||
|
||||
server := http.Server{
|
||||
ReadHeaderTimeout: time.Minute,
|
||||
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
t.Log("got request for", r.URL)
|
||||
// Going to use something _slightly_ exotic so that we can't accidentally get some
|
||||
// default behavior creating a false positive on the test
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
}),
|
||||
BaseContext: func(_ net.Listener) context.Context {
|
||||
return ctx
|
||||
cases := []struct {
|
||||
name string
|
||||
version tunnelsdk.TunnelVersion
|
||||
}{
|
||||
{
|
||||
name: "V1",
|
||||
version: tunnelsdk.TunnelVersion1,
|
||||
},
|
||||
{
|
||||
name: "V2",
|
||||
version: tunnelsdk.TunnelVersion2,
|
||||
},
|
||||
}
|
||||
|
||||
fTunServer := newFakeTunnelServer(t)
|
||||
cfg := fTunServer.config()
|
||||
for _, c := range cases {
|
||||
c := c
|
||||
|
||||
tun, errCh, err := devtunnel.NewWithConfig(ctx, slogtest.Make(t, nil).Leveled(slog.LevelDebug), cfg)
|
||||
require.NoError(t, err)
|
||||
t.Log(tun.URL)
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
go func() {
|
||||
err := server.Serve(tun.Listener)
|
||||
assert.Equal(t, http.ErrServerClosed, err)
|
||||
}()
|
||||
defer func() { _ = server.Close() }()
|
||||
defer func() { tun.Listener.Close() }()
|
||||
ctx, cancelTun := context.WithCancel(context.Background())
|
||||
defer cancelTun()
|
||||
|
||||
require.Eventually(t, func() bool {
|
||||
res, err := fTunServer.requestHTTP()
|
||||
if !assert.NoError(t, err) {
|
||||
return false
|
||||
}
|
||||
defer res.Body.Close()
|
||||
_, _ = io.Copy(io.Discard, res.Body)
|
||||
server := http.Server{
|
||||
ReadHeaderTimeout: time.Minute,
|
||||
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
t.Log("got request for", r.URL)
|
||||
// Going to use something _slightly_ exotic so that we can't
|
||||
// accidentally get some default behavior creating a false
|
||||
// positive on the test
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
}),
|
||||
BaseContext: func(_ net.Listener) context.Context {
|
||||
return ctx
|
||||
},
|
||||
}
|
||||
|
||||
return res.StatusCode == http.StatusAccepted
|
||||
}, testutil.WaitShort, testutil.IntervalSlow)
|
||||
tunServer := newTunnelServer(t)
|
||||
cfg := tunServer.config(t, c.version)
|
||||
|
||||
assert.NoError(t, server.Close())
|
||||
cancelTun()
|
||||
tun, err := devtunnel.NewWithConfig(ctx, slogtest.Make(t, nil).Leveled(slog.LevelDebug), cfg)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, tun.OtherURLs, 1)
|
||||
t.Log(tun.URL, tun.OtherURLs[0])
|
||||
|
||||
select {
|
||||
case <-errCh:
|
||||
case <-time.After(testutil.WaitLong):
|
||||
t.Errorf("tunnel did not close after %s", testutil.WaitLong)
|
||||
hostSplit := strings.SplitN(tun.URL.Host, ".", 2)
|
||||
require.Len(t, hostSplit, 2)
|
||||
require.Equal(t, hostSplit[1], tunServer.api.BaseURL.Host)
|
||||
|
||||
// Verify the hostname using the same logic as the tunnel server.
|
||||
ip1, urls := tunServer.api.WireguardPublicKeyToIPAndURLs(cfg.PublicKey, c.version)
|
||||
require.Len(t, urls, 2)
|
||||
require.Equal(t, urls[0].String(), tun.URL.String())
|
||||
require.Equal(t, urls[1].String(), tun.OtherURLs[0].String())
|
||||
|
||||
ip2, err := tunServer.api.HostnameToWireguardIP(hostSplit[0])
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ip1, ip2)
|
||||
|
||||
// Manually verify the hostname.
|
||||
switch c.version {
|
||||
case tunnelsdk.TunnelVersion1:
|
||||
// The subdomain should be a 32 character hex string.
|
||||
require.Len(t, hostSplit[0], 32)
|
||||
_, err := hex.DecodeString(hostSplit[0])
|
||||
require.NoError(t, err)
|
||||
case tunnelsdk.TunnelVersion2:
|
||||
// The subdomain should be a base32 encoded string containing
|
||||
// 16 bytes once decoded.
|
||||
dec, err := base32.HexEncoding.WithPadding(base32.NoPadding).DecodeString(strings.ToUpper(hostSplit[0]))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, dec, 8)
|
||||
}
|
||||
|
||||
go func() {
|
||||
err := server.Serve(tun.Listener)
|
||||
assert.Equal(t, http.ErrServerClosed, err)
|
||||
}()
|
||||
defer func() { _ = server.Close() }()
|
||||
defer func() { tun.Listener.Close() }()
|
||||
|
||||
require.Eventually(t, func() bool {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, tun.URL.String(), nil)
|
||||
if !assert.NoError(t, err) {
|
||||
return false
|
||||
}
|
||||
res, err := tunServer.requestTunnel(tun, req)
|
||||
if !assert.NoError(t, err) {
|
||||
return false
|
||||
}
|
||||
defer res.Body.Close()
|
||||
_, _ = io.Copy(io.Discard, res.Body)
|
||||
|
||||
return res.StatusCode == http.StatusAccepted
|
||||
}, testutil.WaitShort, testutil.IntervalSlow)
|
||||
|
||||
assert.NoError(t, server.Close())
|
||||
cancelTun()
|
||||
|
||||
select {
|
||||
case <-tun.Wait():
|
||||
case <-time.After(testutil.WaitLong):
|
||||
t.Errorf("tunnel did not close after %s", testutil.WaitLong)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// fakeTunnelServer is a fake version of the real dev tunnel server. It fakes 2 client interactions
|
||||
// that we want to test:
|
||||
// 1. Responding to a POST /tun from the client
|
||||
// 2. Sending an HTTP request down the wireguard connection
|
||||
//
|
||||
// Note that for 2, we don't implement a full proxy that accepts arbitrary requests, we just send
|
||||
// a test request over the Wireguard tunnel to make sure that we can listen. The proxy behavior is
|
||||
// outside of the scope of the dev tunnel client, which is what we are testing here.
|
||||
type fakeTunnelServer struct {
|
||||
t *testing.T
|
||||
pub device.NoisePublicKey
|
||||
priv device.NoisePrivateKey
|
||||
tnet *netstack.Net
|
||||
device *device.Device
|
||||
clients int
|
||||
server *httptest.Server
|
||||
}
|
||||
|
||||
func newFakeTunnelServer(t *testing.T) *fakeTunnelServer {
|
||||
func freeUDPPort(t *testing.T) uint16 {
|
||||
t.Helper()
|
||||
|
||||
priv, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(t, err)
|
||||
privBytes := [32]byte(priv)
|
||||
pub := priv.PublicKey()
|
||||
pubBytes := [32]byte(pub)
|
||||
tun, tnet, err := netstack.CreateNetTUN(
|
||||
[]netip.Addr{serverIP},
|
||||
[]netip.Addr{dnsIP},
|
||||
1280,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
ctx := context.Background()
|
||||
slogger := slogtest.Make(t, nil).Leveled(slog.LevelDebug).Named("server")
|
||||
logger := &device.Logger{
|
||||
Verbosef: slog.Stdlib(ctx, slogger, slog.LevelDebug).Printf,
|
||||
Errorf: slog.Stdlib(ctx, slogger, slog.LevelError).Printf,
|
||||
}
|
||||
dev := device.NewDevice(tun, conn.NewDefaultBind(), logger)
|
||||
t.Cleanup(func() {
|
||||
dev.RemoveAllPeers()
|
||||
dev.Close()
|
||||
slogger.Debug(ctx, "dev.Close()")
|
||||
l, err := net.ListenUDP("udp", &net.UDPAddr{
|
||||
IP: net.ParseIP("127.0.0.1"),
|
||||
Port: 0,
|
||||
})
|
||||
err = dev.IpcSet(fmt.Sprintf(`private_key=%s
|
||||
listen_port=%d`,
|
||||
hex.EncodeToString(privBytes[:]),
|
||||
wgPort,
|
||||
))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, err, "listen on random UDP port")
|
||||
|
||||
err = dev.Up()
|
||||
require.NoError(t, err)
|
||||
_, port, err := net.SplitHostPort(l.LocalAddr().String())
|
||||
require.NoError(t, err, "split host port")
|
||||
|
||||
server := newFakeTunnelHTTPSServer(t, pubBytes)
|
||||
portUint, err := strconv.ParseUint(port, 10, 16)
|
||||
require.NoError(t, err, "parse port")
|
||||
|
||||
return &fakeTunnelServer{
|
||||
t: t,
|
||||
pub: device.NoisePublicKey(pub),
|
||||
priv: device.NoisePrivateKey(priv),
|
||||
tnet: tnet,
|
||||
device: dev,
|
||||
server: server,
|
||||
}
|
||||
// This is prone to races, but since we have to tell wireguard to create the
|
||||
// listener and can't pass in a net.Listener, we have to do this.
|
||||
err = l.Close()
|
||||
require.NoError(t, err, "close UDP listener")
|
||||
|
||||
return uint16(portUint)
|
||||
}
|
||||
|
||||
func newFakeTunnelHTTPSServer(t *testing.T, pubBytes [32]byte) *httptest.Server {
|
||||
handler := http.NewServeMux()
|
||||
handler.HandleFunc("/tun", func(writer http.ResponseWriter, request *http.Request) {
|
||||
assert.Equal(t, "POST", request.Method)
|
||||
type tunnelServer struct {
|
||||
api *tunneld.API
|
||||
|
||||
resp := devtunnel.ServerResponse{
|
||||
Hostname: fmt.Sprintf("[%s]", serverIP.String()),
|
||||
ServerIP: serverIP,
|
||||
ServerPublicKey: hex.EncodeToString(pubBytes[:]),
|
||||
ClientIP: clientIP,
|
||||
server *httptest.Server
|
||||
}
|
||||
|
||||
func newTunnelServer(t *testing.T) *tunnelServer {
|
||||
var handler http.Handler
|
||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if handler != nil {
|
||||
handler.ServeHTTP(w, r)
|
||||
}
|
||||
b, err := json.Marshal(&resp)
|
||||
assert.NoError(t, err)
|
||||
writer.WriteHeader(200)
|
||||
_, err = writer.Write(b)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
|
||||
server := httptest.NewTLSServer(handler)
|
||||
w.WriteHeader(http.StatusBadGateway)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
baseURLParsed, err := url.Parse(srv.URL)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https", baseURLParsed.Scheme)
|
||||
baseURLParsed.Host = net.JoinHostPort("tunnel.coder.com", baseURLParsed.Port())
|
||||
|
||||
wireguardPort := freeUDPPort(t)
|
||||
|
||||
key, err := tunnelsdk.GeneratePrivateKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
options := &tunneld.Options{
|
||||
BaseURL: baseURLParsed,
|
||||
WireguardEndpoint: fmt.Sprintf("127.0.0.1:%d", wireguardPort),
|
||||
WireguardPort: wireguardPort,
|
||||
WireguardKey: key,
|
||||
WireguardMTU: tunneld.DefaultWireguardMTU,
|
||||
WireguardServerIP: tunneld.DefaultWireguardServerIP,
|
||||
WireguardNetworkPrefix: tunneld.DefaultWireguardNetworkPrefix,
|
||||
}
|
||||
|
||||
td, err := tunneld.New(options)
|
||||
require.NoError(t, err)
|
||||
handler = td.Router()
|
||||
t.Cleanup(func() {
|
||||
server.Close()
|
||||
_ = td.Close()
|
||||
})
|
||||
return server
|
||||
}
|
||||
|
||||
func (f *fakeTunnelServer) config() devtunnel.Config {
|
||||
priv, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(f.t, err)
|
||||
pub := priv.PublicKey()
|
||||
f.clients++
|
||||
assert.Equal(f.t, 1, f.clients) // only allow one client as we hardcode the address
|
||||
|
||||
err = f.device.IpcSet(fmt.Sprintf(`public_key=%x
|
||||
allowed_ip=%s/128`,
|
||||
pub[:],
|
||||
clientIP.String(),
|
||||
))
|
||||
require.NoError(f.t, err)
|
||||
return devtunnel.Config{
|
||||
Version: 1,
|
||||
PrivateKey: device.NoisePrivateKey(priv),
|
||||
PublicKey: device.NoisePublicKey(pub),
|
||||
Tunnel: devtunnel.Node{
|
||||
HostnameHTTPS: strings.TrimPrefix(f.server.URL, "https://"),
|
||||
HostnameWireguard: "localhost",
|
||||
WireguardPort: wgPort,
|
||||
},
|
||||
HTTPClient: f.server.Client(),
|
||||
return &tunnelServer{
|
||||
api: td,
|
||||
server: srv,
|
||||
}
|
||||
}
|
||||
|
||||
func (f *fakeTunnelServer) requestHTTP() (*http.Response, error) {
|
||||
func (s *tunnelServer) client() *http.Client {
|
||||
transport := &http.Transport{
|
||||
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
f.t.Log("Dial", network, addr)
|
||||
nc, err := f.tnet.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(clientIP, 8090))
|
||||
assert.NoError(f.t, err)
|
||||
return nc, err
|
||||
return (&net.Dialer{}).DialContext(ctx, "tcp", s.server.Listener.Addr().String())
|
||||
},
|
||||
TLSClientConfig: &tls.Config{
|
||||
//nolint:gosec
|
||||
InsecureSkipVerify: true,
|
||||
},
|
||||
}
|
||||
client := &http.Client{
|
||||
return &http.Client{
|
||||
Transport: transport,
|
||||
Timeout: testutil.WaitLong,
|
||||
}
|
||||
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, fmt.Sprintf("http://[%s]:8090", clientIP), nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
func (s *tunnelServer) config(t *testing.T, version tunnelsdk.TunnelVersion) devtunnel.Config {
|
||||
priv, err := tunnelsdk.GeneratePrivateKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
privNoise, err := priv.NoisePrivateKey()
|
||||
require.NoError(t, err)
|
||||
pubNoise := priv.NoisePublicKey()
|
||||
|
||||
if version == 0 {
|
||||
version = tunnelsdk.TunnelVersionLatest
|
||||
}
|
||||
return client.Do(req)
|
||||
|
||||
return devtunnel.Config{
|
||||
Version: version,
|
||||
PrivateKey: privNoise,
|
||||
PublicKey: pubNoise,
|
||||
Tunnel: devtunnel.Node{
|
||||
RegionID: 0,
|
||||
ID: 1,
|
||||
HostnameHTTPS: s.api.BaseURL.Host,
|
||||
},
|
||||
HTTPClient: s.client(),
|
||||
}
|
||||
}
|
||||
|
||||
// requestTunnel performs the given request against the tunnel. The Host header
|
||||
// will be set to the tunnel's hostname.
|
||||
func (s *tunnelServer) requestTunnel(tunnel *tunnelsdk.Tunnel, req *http.Request) (*http.Response, error) {
|
||||
req.URL.Scheme = "https"
|
||||
req.URL.Host = tunnel.URL.Host
|
||||
req.Host = tunnel.URL.Host
|
||||
return s.client().Do(req)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user