package netutil import ( "bufio" "context" "errors" "fmt" "net" "net/http" "net/netip" "slices" "testing" "time" ) type mockResolver struct { ips map[string][]netip.Addr err map[string]error } func (m *mockResolver) LookupIP(_ context.Context, host string) ([]netip.Addr, error) { if err, ok := m.err[host]; ok { return nil, err } if ips, ok := m.ips[host]; ok { return ips, nil } return nil, errors.New("nxdomain") } type mockDialer struct { dialedAddr string attempts []string failOn map[string]error } func (m *mockDialer) DialContext(_ context.Context, _, address string) (net.Conn, error) { m.dialedAddr = address m.attempts = append(m.attempts, address) if err, ok := m.failOn[address]; ok { return nil, err } server, client := net.Pipe() go func() { _ = server.Close() }() return client, nil } func TestProxyRejectsNonConnect(t *testing.T) { p := NewSafeConnectProxy(nil, nil) addr, err := p.Start() if err != nil { t.Fatal(err) } defer p.Close() resp, err := http.Get("http://" + addr + "/foo") if err != nil { t.Fatal(err) } defer resp.Body.Close() if resp.StatusCode != http.StatusMethodNotAllowed { t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusMethodNotAllowed) } } func TestProxyRejectsBadPort(t *testing.T) { p := NewSafeConnectProxy(nil, nil) addr, err := p.Start() if err != nil { t.Fatal(err) } defer p.Close() conn, err := net.Dial("tcp", addr) if err != nil { t.Fatal(err) } defer conn.Close() fmt.Fprintf(conn, "CONNECT github.com:80 HTTP/1.1\r\nHost: github.com:80\r\n\r\n") resp, err := http.ReadResponse(bufio.NewReader(conn), nil) if err != nil { t.Fatal(err) } if resp.StatusCode != http.StatusForbidden { t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusForbidden) } } func TestProxyRejectsDirectPrivateIP(t *testing.T) { p := NewSafeConnectProxy(nil, nil) addr, err := p.Start() if err != nil { t.Fatal(err) } defer p.Close() for _, target := range []string{"127.0.0.1:443", "10.0.0.1:443", "192.168.1.5:443", "169.254.169.254:443", "[::1]:443", "[2002:7f00:1::]:443"} { conn, err := net.Dial("tcp", addr) if err != nil { t.Fatal(err) } fmt.Fprintf(conn, "CONNECT %s HTTP/1.1\r\nHost: %s\r\n\r\n", target, target) resp, err := http.ReadResponse(bufio.NewReader(conn), nil) if err != nil { conn.Close() t.Fatal(err) } conn.Close() if resp.StatusCode != http.StatusForbidden { t.Fatalf("target %s returned status %d, want 403 Forbidden", target, resp.StatusCode) } } } func TestProxyRejectsMixedDNSAnswer(t *testing.T) { res := &mockResolver{ ips: map[string][]netip.Addr{ "mixed.example.com": { netip.MustParseAddr("140.82.112.4"), // Safe public IP netip.MustParseAddr("127.0.0.1"), }, }, } p := NewSafeConnectProxy(res, nil) addr, err := p.Start() if err != nil { t.Fatal(err) } defer p.Close() conn, err := net.Dial("tcp", addr) if err != nil { t.Fatal(err) } defer conn.Close() fmt.Fprintf(conn, "CONNECT mixed.example.com:443 HTTP/1.1\r\nHost: mixed.example.com:443\r\n\r\n") resp, err := http.ReadResponse(bufio.NewReader(conn), nil) if err != nil { t.Fatal(err) } if resp.StatusCode != http.StatusForbidden { t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusForbidden) } } func TestProxyPermitsSafeDestinationAndDialsPinnedIP(t *testing.T) { safeIP := netip.MustParseAddr("140.82.112.4") res := &mockResolver{ ips: map[string][]netip.Addr{ "github.com": {safeIP}, }, } dialer := &mockDialer{} p := NewSafeConnectProxy(res, dialer) addr, err := p.Start() if err != nil { t.Fatal(err) } defer p.Close() conn, err := net.Dial("tcp", addr) if err != nil { t.Fatal(err) } defer conn.Close() fmt.Fprintf(conn, "CONNECT github.com:443 HTTP/1.1\r\nHost: github.com:443\r\n\r\n") resp, err := http.ReadResponse(bufio.NewReader(conn), nil) if err != nil { t.Fatal(err) } if resp.StatusCode != http.StatusOK { t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusOK) } expectedDial := safeIP.String() + ":443" if dialer.dialedAddr != expectedDial { t.Fatalf("dialed %q, want %q", dialer.dialedAddr, expectedDial) } } func TestProxyTriesEveryResolvedAddress(t *testing.T) { unreachable := netip.MustParseAddr("140.82.112.4") reachable := netip.MustParseAddr("140.82.113.4") res := &mockResolver{ ips: map[string][]netip.Addr{"github.com": {unreachable, reachable}}, } dialer := &mockDialer{ failOn: map[string]error{unreachable.String() + ":443": errors.New("network is unreachable")}, } p := NewSafeConnectProxy(res, dialer) addr, err := p.Start() if err != nil { t.Fatal(err) } defer p.Close() conn, err := net.Dial("tcp", addr) if err != nil { t.Fatal(err) } defer conn.Close() fmt.Fprintf(conn, "CONNECT github.com:443 HTTP/1.1\r\nHost: github.com:443\r\n\r\n") resp, err := http.ReadResponse(bufio.NewReader(conn), nil) if err != nil { t.Fatal(err) } if resp.StatusCode != http.StatusOK { t.Fatalf("status = %d, want %d: a v4-only network is not a refusal", resp.StatusCode, http.StatusOK) } want := []string{unreachable.String() + ":443", reachable.String() + ":443"} if !slices.Equal(dialer.attempts, want) { t.Fatalf("dialled %v, want %v", dialer.attempts, want) } } func TestProxyReleasesATunnelThatStopsTalking(t *testing.T) { old := relayDrainGrace relayDrainGrace = 100 * time.Millisecond defer func() { relayDrainGrace = old }() safeIP := netip.MustParseAddr("140.82.112.4") res := &mockResolver{ips: map[string][]netip.Addr{"github.com": {safeIP}}} serverConn, peerConn := net.Pipe() defer peerConn.Close() dialer := &pipeDialer{conn: serverConn} p := NewSafeConnectProxy(res, dialer) addr, err := p.Start() if err != nil { t.Fatal(err) } defer p.Close() conn, err := net.Dial("tcp", addr) if err != nil { t.Fatal(err) } fmt.Fprintf(conn, "CONNECT github.com:443 HTTP/1.1\r\nHost: github.com:443\r\n\r\n") if _, err := http.ReadResponse(bufio.NewReader(conn), nil); err != nil { t.Fatal(err) } conn.Close() deadline := time.Now().Add(3 * time.Second) for time.Now().Before(deadline) { p.mu.Lock() open := len(p.active) p.mu.Unlock() if open == 0 { return } time.Sleep(20 * time.Millisecond) } t.Fatal("the tunnel outlived the relay: a half-closed pair is never reaped") } type pipeDialer struct{ conn net.Conn } func (d *pipeDialer) DialContext(context.Context, string, string) (net.Conn, error) { return d.conn, nil }