Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 10 additions & 4 deletions config.go
Original file line number Diff line number Diff line change
Expand Up @@ -319,6 +319,9 @@ func ParseInterface(cfg *ini.File, device *DeviceConfig) error {
if len(checkAlive) == 0 {
return errors.New("CheckAliveInterval is only valid when CheckAlive is set")
}
if value <= 0 {
return errors.New("CheckAliveInterval should be greater than zero")
}

device.CheckAliveInterval = value
}
Expand Down Expand Up @@ -558,7 +561,7 @@ func parseResolveConfig(section *ini.Section) (*ResolveConfig, error) {

resolvStrategy, _ := parseString(section, "ResolveStrategy")
config.ResolveStrategy = resolvStrategy

return config, nil
}

Expand All @@ -583,6 +586,9 @@ func parseUDPProxyTunnelConfig(section *ini.Section) (RoutineSpawner, error) {
if err != nil {
return nil, err
}
if timeoutVal < 0 {
return nil, errors.New("InactivityTimeout should not be negative")
}
inactivityTimeout = timeoutVal
}
config.InactivityTimeout = inactivityTimeout
Expand Down Expand Up @@ -694,9 +700,9 @@ func ParseConfig(path string) (*Configuration, error) {
resolve, err = parseResolveConfig(resolveSection)
if err != nil {
return nil, err
}
}
}
}

err = parseRoutinesConfig(&routinesSpawners, cfg, "UDPProxyTunnel", parseUDPProxyTunnelConfig)
if err != nil {
return nil, err
Expand Down
39 changes: 39 additions & 0 deletions config_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -85,3 +85,42 @@ Endpoint = 192.200.144.22:51820`
t.Fatal(err)
}
}

func TestCheckAliveIntervalMustBePositive(t *testing.T) {
const config = `
[Interface]
PrivateKey = LAr1aNSNF9d0MjwUgAVC4020T0N/E5NUtqVv5EnsSz0=
Address = 10.5.0.2
CheckAlive = 1.1.1.1
CheckAliveInterval = 0

[Peer]
PublicKey = e8LKAc+f9xEzq9Ar7+MfKRrs+gZ/4yzvpRJLRJ/VJ1w=`

iniData, err := loadIniConfig(config)
if err != nil {
t.Fatal(err)
}

var cfg DeviceConfig
if err := ParseInterface(iniData, &cfg); err == nil {
t.Fatal("expected non-positive CheckAliveInterval to be rejected")
}
}

func TestInactivityTimeoutCannotBeNegative(t *testing.T) {
const config = `
[UDPProxyTunnel]
BindAddress = 127.0.0.1:25346
Target = 1.1.1.1:53
InactivityTimeout = -1`

iniData, err := loadIniConfig(config)
if err != nil {
t.Fatal(err)
}

if _, err := parseUDPProxyTunnelConfig(iniData.Section("UDPProxyTunnel")); err == nil {
t.Fatal("expected negative InactivityTimeout to be rejected")
}
}
11 changes: 11 additions & 0 deletions http.go
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,12 @@ func (s *HTTPServer) handle(req *http.Request) (peer net.Conn, err error) {
}

func (s *HTTPServer) serve(conn net.Conn) {
var handled bool
defer func() {
if !handled {
_ = conn.Close()
}
}()
var rd = bufio.NewReader(conn)
req, err := http.ReadRequest(rd)
if err != nil {
Expand Down Expand Up @@ -124,6 +130,9 @@ func (s *HTTPServer) serve(conn net.Conn) {
return
}
if err != nil {
if peer != nil {
_ = peer.Close()
}
log.Printf("dial proxy failed: %s\n", err)
return
}
Expand All @@ -132,6 +141,8 @@ func (s *HTTPServer) serve(conn net.Conn) {
return
}

handled = true

go func() {
defer func() { _ = conn.Close() }()
defer func() { _ = peer.Close() }()
Expand Down
101 changes: 101 additions & 0 deletions http_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
package wireproxy

import (
"bufio"
"io"
"net"
"net/http"
"net/http/httptest"
"testing"
)

func TestHTTPServerServeGet(t *testing.T) {
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte("ok from upstream"))
}))
defer ts.Close()

server := &HTTPServer{
dial: func(network, address string) (net.Conn, error) {
return net.Dial(network, ts.Listener.Addr().String())
},
}

listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer func() { _ = listener.Close() }()

go func() {
conn, err := listener.Accept()
if err != nil {
return
}
server.serve(conn)
}()

clientConn, err := net.Dial("tcp", listener.Addr().String())
if err != nil {
t.Fatal(err)
}
defer func() { _ = clientConn.Close() }()

req, err := http.NewRequest(http.MethodGet, "http://"+ts.Listener.Addr().String()+"/test", nil)
if err != nil {
t.Fatal(err)
}

if err := req.Write(clientConn); err != nil {
t.Fatal(err)
}

resp, err := http.ReadResponse(bufio.NewReader(clientConn), req)
if err != nil {
t.Fatalf("failed to read response: %v", err)
}
defer func() { _ = resp.Body.Close() }()

body, err := io.ReadAll(resp.Body)
if err != nil {
t.Fatalf("failed to read body: %v", err)
}
if string(body) != "ok from upstream" {
t.Fatalf("unexpected body: %q", string(body))
}
}

func TestHTTPServerServeFailureClosesConn(t *testing.T) {
server := &HTTPServer{}

listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer func() { _ = listener.Close() }()

go func() {
conn, err := listener.Accept()
if err != nil {
return
}
server.serve(conn)
}()

clientConn, err := net.Dial("tcp", listener.Addr().String())
if err != nil {
t.Fatal(err)
}
defer func() { _ = clientConn.Close() }()

// Send invalid HTTP request
if _, err := clientConn.Write([]byte("INVALID\r\n\r\n")); err != nil {
t.Fatal(err)
}

buf := make([]byte, 1)
_, err = clientConn.Read(buf)
if err != io.EOF {
t.Fatalf("expected EOF when server closes connection, got %v", err)
}
}
28 changes: 28 additions & 0 deletions leak_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
package wireproxy

import (
"net"
"testing"
)

func TestUDPSessionCloseIsIdempotent(t *testing.T) {
local, remote := net.Pipe()
session := &udpSession{
remoteConn: local,
closeChan: make(chan struct{}),
}

session.close()
session.close()

select {
case <-session.closeChan:
default:
t.Fatal("session close channel was not closed")
}

if _, err := remote.Write([]byte("data")); err == nil {
t.Fatal("remote connection remained open after session close")
}
_ = remote.Close()
}
60 changes: 45 additions & 15 deletions routine.go
Original file line number Diff line number Diff line change
Expand Up @@ -277,6 +277,7 @@ func connForward(from io.ReadWriteCloser, to io.ReadWriteCloser) {
func tcpClientForward(vt *VirtualTun, raddr *addressPort, conn net.Conn) {
target, err := vt.resolveToAddrPort(raddr)
if err != nil {
_ = conn.Close()
errorLogger.Printf("TCP Server Tunnel to %s: %s\n", target, err.Error())
return
}
Expand All @@ -285,6 +286,7 @@ func tcpClientForward(vt *VirtualTun, raddr *addressPort, conn net.Conn) {

sconn, err := vt.Tnet.DialTCP(tcpAddr)
if err != nil {
_ = conn.Close()
errorLogger.Printf("TCP Client Tunnel to %s: %s\n", target, err.Error())
return
}
Expand Down Expand Up @@ -347,6 +349,7 @@ func (conf *STDIOTunnelConfig) SpawnRoutine(vt *VirtualTun) {
func tcpServerForward(vt *VirtualTun, raddr *addressPort, conn net.Conn) {
target, err := vt.resolveToAddrPort(raddr)
if err != nil {
_ = conn.Close()
errorLogger.Printf("TCP Server Tunnel to %s: %s\n", target, err.Error())
return
}
Expand All @@ -355,6 +358,7 @@ func tcpServerForward(vt *VirtualTun, raddr *addressPort, conn net.Conn) {

sconn, err := net.DialTCP("tcp", nil, tcpAddr)
if err != nil {
_ = conn.Close()
errorLogger.Printf("TCP Server Tunnel to %s: %s\n", target, err.Error())
return
}
Expand Down Expand Up @@ -409,15 +413,22 @@ func (d VirtualTun) ServeHTTP(w http.ResponseWriter, r *http.Request) {
log.Printf("Health metric request: %s\n", r.URL.Path)
switch path.Clean(r.URL.Path) {
case "/readyz":
body, err := json.Marshal(d.PingRecord)
d.PingRecordLock.Lock()
records := make(map[string]uint64, len(d.PingRecord))
for addr, record := range d.PingRecord {
records[addr] = record
}
d.PingRecordLock.Unlock()

body, err := json.Marshal(records)
if err != nil {
errorLogger.Printf("Failed to get device metrics: %s\n", err.Error())
w.WriteHeader(http.StatusInternalServerError)
return
}

status := http.StatusOK
for _, record := range d.PingRecord {
for _, record := range records {
lastPong := time.Unix(int64(record), 0)
// +2 seconds to account for the time it takes to ping the IP
if time.Since(lastPong) > time.Duration(d.Conf.CheckAliveInterval+2)*time.Second {
Expand Down Expand Up @@ -466,9 +477,18 @@ func (d VirtualTun) pingIPs() {
errorLogger.Printf("Failed to ping %s: %s\n", addr, err.Error())
continue
}
if !addr.Is4() && !addr.Is6() {
errorLogger.Printf("Failed to ping %s: invalid address: %s\n", addr, addr.String())
_ = socket.Close()
continue
}

data := make([]byte, 16)
_, _ = srand.Read(data)
if _, err := srand.Read(data); err != nil {
errorLogger.Printf("Failed to generate ping data for %s: %s\n", addr, err.Error())
_ = socket.Close()
continue
}

requestPing := icmp.Echo{
Seq: rand.Intn(1 << 16),
Expand All @@ -477,30 +497,38 @@ func (d VirtualTun) pingIPs() {

var icmpBytes []byte
if addr.Is4() {
icmpBytes, _ = (&icmp.Message{Type: ipv4.ICMPTypeEcho, Code: 0, Body: &requestPing}).Marshal(nil)
icmpBytes, err = (&icmp.Message{Type: ipv4.ICMPTypeEcho, Code: 0, Body: &requestPing}).Marshal(nil)
} else if addr.Is6() {
icmpBytes, _ = (&icmp.Message{Type: ipv6.ICMPTypeEchoRequest, Code: 0, Body: &requestPing}).Marshal(nil)
} else {
errorLogger.Printf("Failed to ping %s: invalid address: %s\n", addr, addr.String())
icmpBytes, err = (&icmp.Message{Type: ipv6.ICMPTypeEchoRequest, Code: 0, Body: &requestPing}).Marshal(nil)
}
if err != nil {
errorLogger.Printf("Failed to marshal ping request for %s: %s\n", addr, err.Error())
_ = socket.Close()
continue
}

_ = socket.SetReadDeadline(time.Now().Add(time.Duration(d.Conf.CheckAliveInterval) * time.Second))
if err := socket.SetReadDeadline(time.Now().Add(time.Duration(d.Conf.CheckAliveInterval) * time.Second)); err != nil {
errorLogger.Printf("Failed to set ping deadline for %s: %s\n", addr, err.Error())
_ = socket.Close()
continue
}
_, err = socket.Write(icmpBytes)
if err != nil {
errorLogger.Printf("Failed to ping %s: %s\n", addr, err.Error())
_ = socket.Close()
continue
}

addr := addr
go func() {
n, err := socket.Read(icmpBytes[:])
go func(addr netip.Addr, socket net.Conn, requestPing icmp.Echo) {
defer func() { _ = socket.Close() }()
readBytes := make([]byte, 1500)
n, err := socket.Read(readBytes)
if err != nil {
errorLogger.Printf("Failed to read ping response from %s: %s\n", addr, err.Error())
return
}

replyPacket, err := icmp.ParseMessage(1, icmpBytes[:n])
replyPacket, err := icmp.ParseMessage(1, readBytes[:n])
if err != nil {
errorLogger.Printf("Failed to parse ping response from %s: %s\n", addr, err.Error())
return
Expand All @@ -524,6 +552,10 @@ func (d VirtualTun) pingIPs() {
errorLogger.Printf("Failed to parse ping response from %s: invalid reply type: %s\n", addr, replyPacket.Type)
return
}
if len(replyPing.Data) < 4 {
errorLogger.Printf("Failed to parse ping response from %s: reply too short\n", addr)
return
}

seq := binary.BigEndian.Uint16(replyPing.Data[2:4])
pongBody := replyPing.Data[4:]
Expand All @@ -536,9 +568,7 @@ func (d VirtualTun) pingIPs() {
d.PingRecordLock.Lock()
d.PingRecord[addr.String()] = uint64(time.Now().Unix())
d.PingRecordLock.Unlock()

defer func() { _ = socket.Close() }()
}()
}(addr, socket, requestPing)
}
}

Expand Down
Loading
Loading