From d15280e94cb18bebb017ba765baa1ee7ab7e9191 Mon Sep 17 00:00:00 2001 From: NikoCat233 <139348239+NikoCat233@users.noreply.github.com> Date: Wed, 2 Sep 2026 10:58:56 +0800 Subject: [PATCH 1/3] fix: close leaked connections and synchronize UDP sessions --- config.go | 14 ++++++++---- config_test.go | 39 ++++++++++++++++++++++++++++++++ http.go | 1 + leak_test.go | 28 +++++++++++++++++++++++ routine.go | 60 +++++++++++++++++++++++++++++++++++++------------- udp_proxy.go | 47 +++++++++++++++++++++++++-------------- wireguard.go | 2 ++ 7 files changed, 156 insertions(+), 35 deletions(-) create mode 100644 leak_test.go diff --git a/config.go b/config.go index a37957d3..9bf53ff4 100644 --- a/config.go +++ b/config.go @@ -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 } @@ -558,7 +561,7 @@ func parseResolveConfig(section *ini.Section) (*ResolveConfig, error) { resolvStrategy, _ := parseString(section, "ResolveStrategy") config.ResolveStrategy = resolvStrategy - + return config, nil } @@ -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 @@ -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 diff --git a/config_test.go b/config_test.go index 948fbf83..be1d36a3 100644 --- a/config_test.go +++ b/config_test.go @@ -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") + } +} diff --git a/http.go b/http.go index fc80b1e1..e9d5521d 100644 --- a/http.go +++ b/http.go @@ -94,6 +94,7 @@ func (s *HTTPServer) handle(req *http.Request) (peer net.Conn, err error) { } func (s *HTTPServer) serve(conn net.Conn) { + defer func() { _ = conn.Close() }() var rd = bufio.NewReader(conn) req, err := http.ReadRequest(rd) if err != nil { diff --git a/leak_test.go b/leak_test.go new file mode 100644 index 00000000..b27fdf6e --- /dev/null +++ b/leak_test.go @@ -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() +} diff --git a/routine.go b/routine.go index 1b3b8850..1ccd2b80 100644 --- a/routine.go +++ b/routine.go @@ -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 } @@ -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 } @@ -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 } @@ -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 } @@ -409,7 +413,14 @@ 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) @@ -417,7 +428,7 @@ func (d VirtualTun) ServeHTTP(w http.ResponseWriter, r *http.Request) { } 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 { @@ -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), @@ -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 @@ -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:] @@ -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) } } diff --git a/udp_proxy.go b/udp_proxy.go index c6e34c48..5cb3ae7b 100644 --- a/udp_proxy.go +++ b/udp_proxy.go @@ -15,6 +15,27 @@ type udpSession struct { lastActive time.Time closeChan chan struct{} inactivityDur time.Duration + closeOnce sync.Once + activityMu sync.Mutex +} + +func (s *udpSession) touch() { + s.activityMu.Lock() + s.lastActive = time.Now() + s.activityMu.Unlock() +} + +func (s *udpSession) inactive(now time.Time) bool { + s.activityMu.Lock() + defer s.activityMu.Unlock() + return now.Sub(s.lastActive) >= s.inactivityDur +} + +func (s *udpSession) close() { + s.closeOnce.Do(func() { + close(s.closeChan) + _ = s.remoteConn.Close() + }) } // SpawnRoutine implements the RoutineSpawner interface. @@ -36,18 +57,10 @@ func (conf *UDPProxyTunnelConfig) SpawnRoutine(vt *VirtualTun) { sessions := make(map[string]*udpSession) var sessionMu sync.Mutex - closeSessionChan := func(sess *udpSession) { - select { - case <-sess.closeChan: - default: - close(sess.closeChan) - } - } - removeSession := func(src string, sess *udpSession) { sessionMu.Lock() if current, ok := sessions[src]; ok && current == sess { - closeSessionChan(current) + current.close() delete(sessions, src) } sessionMu.Unlock() @@ -62,9 +75,9 @@ func (conf *UDPProxyTunnelConfig) SpawnRoutine(vt *VirtualTun) { now := time.Now() sessionMu.Lock() for key, sess := range sessions { - if now.Sub(sess.lastActive) >= inactivityDur { + if sess.inactive(now) { log.Printf("UDPProxyTunnel: closing inactive session for %s", key) - closeSessionChan(sess) + sess.close() delete(sessions, key) } } @@ -80,7 +93,7 @@ func (conf *UDPProxyTunnelConfig) SpawnRoutine(vt *VirtualTun) { // return if session already exists if s, ok := sessions[srcAddr]; ok { - s.lastActive = time.Now() + s.touch() return s, nil } @@ -120,10 +133,10 @@ func (conf *UDPProxyTunnelConfig) SpawnRoutine(vt *VirtualTun) { continue } - s.lastActive = time.Now() _, err = s.remoteConn.Write(buf[:n]) if err != nil { errorLogger.Printf("UDPProxyTunnel: could not write to remote (%s): %v", conf.Target, err) + removeSession(srcKey, s) } } }() @@ -134,7 +147,6 @@ func (conf *UDPProxyTunnelConfig) SpawnRoutine(vt *VirtualTun) { func (conf *UDPProxyTunnelConfig) handleRemoteToLocal(listener *net.UDPConn, srcAddr string, s *udpSession, removeSession func(string, *udpSession)) { defer func() { removeSession(srcAddr, s) - _ = s.remoteConn.Close() }() buf := make([]byte, 64*1024) @@ -145,7 +157,10 @@ func (conf *UDPProxyTunnelConfig) handleRemoteToLocal(listener *net.UDPConn, src default: } - _ = s.remoteConn.SetReadDeadline(time.Now().Add(5 * time.Second)) + if err := s.remoteConn.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil { + errorLogger.Printf("UDPProxyTunnel: could not set remote read deadline: %v", err) + return + } n, err := s.remoteConn.Read(buf) if err != nil { // If a timeout or temporary error, continue to see if the session is closed @@ -161,7 +176,7 @@ func (conf *UDPProxyTunnelConfig) handleRemoteToLocal(listener *net.UDPConn, src return } - s.lastActive = time.Now() + s.touch() dstUDPAddr, err := net.ResolveUDPAddr("udp", srcAddr) if err != nil { diff --git a/wireguard.go b/wireguard.go index 3da5fa67..0606eb88 100644 --- a/wireguard.go +++ b/wireguard.go @@ -74,11 +74,13 @@ func StartWireguard(conf *Configuration, logLevel int) (*VirtualTun, error) { dev := device.NewDevice(tun, conn.NewDefaultBind(), device.NewLogger(logLevel, "")) err = dev.IpcSet(setting.IpcRequest) if err != nil { + dev.Close() return nil, err } err = dev.Up() if err != nil { + dev.Close() return nil, err } From de966815af7cbf72ec72497ecc29c64444623dc5 Mon Sep 17 00:00:00 2001 From: Wind Date: Wed, 2 Sep 2026 22:48:05 +0200 Subject: [PATCH 2/3] fix(http): do not prematurely close connection in HTTPServer.serve --- http.go | 12 +++++- http_test.go | 101 +++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 112 insertions(+), 1 deletion(-) create mode 100644 http_test.go diff --git a/http.go b/http.go index e9d5521d..110bf9ef 100644 --- a/http.go +++ b/http.go @@ -94,7 +94,12 @@ func (s *HTTPServer) handle(req *http.Request) (peer net.Conn, err error) { } func (s *HTTPServer) serve(conn net.Conn) { - defer func() { _ = conn.Close() }() + var handled bool + defer func() { + if !handled { + _ = conn.Close() + } + }() var rd = bufio.NewReader(conn) req, err := http.ReadRequest(rd) if err != nil { @@ -125,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 } @@ -133,6 +141,8 @@ func (s *HTTPServer) serve(conn net.Conn) { return } + handled = true + go func() { defer func() { _ = conn.Close() }() defer func() { _ = peer.Close() }() diff --git a/http_test.go b/http_test.go new file mode 100644 index 00000000..66d61fec --- /dev/null +++ b/http_test.go @@ -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 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 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 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 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 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) + } +} From 2b6b3a14347af102f739b208178e5c80933cbf32 Mon Sep 17 00:00:00 2001 From: Wind Date: Wed, 2 Sep 2026 22:49:52 +0200 Subject: [PATCH 3/3] test: ignore Close errors in defer statements for errcheck --- http_test.go | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/http_test.go b/http_test.go index 66d61fec..97ccc02d 100644 --- a/http_test.go +++ b/http_test.go @@ -25,7 +25,7 @@ func TestHTTPServerServeGet(t *testing.T) { if err != nil { t.Fatal(err) } - defer listener.Close() + defer func() { _ = listener.Close() }() go func() { conn, err := listener.Accept() @@ -39,7 +39,7 @@ func TestHTTPServerServeGet(t *testing.T) { if err != nil { t.Fatal(err) } - defer clientConn.Close() + defer func() { _ = clientConn.Close() }() req, err := http.NewRequest(http.MethodGet, "http://"+ts.Listener.Addr().String()+"/test", nil) if err != nil { @@ -54,7 +54,7 @@ func TestHTTPServerServeGet(t *testing.T) { if err != nil { t.Fatalf("failed to read response: %v", err) } - defer resp.Body.Close() + defer func() { _ = resp.Body.Close() }() body, err := io.ReadAll(resp.Body) if err != nil { @@ -72,7 +72,7 @@ func TestHTTPServerServeFailureClosesConn(t *testing.T) { if err != nil { t.Fatal(err) } - defer listener.Close() + defer func() { _ = listener.Close() }() go func() { conn, err := listener.Accept() @@ -86,7 +86,7 @@ func TestHTTPServerServeFailureClosesConn(t *testing.T) { if err != nil { t.Fatal(err) } - defer clientConn.Close() + defer func() { _ = clientConn.Close() }() // Send invalid HTTP request if _, err := clientConn.Write([]byte("INVALID\r\n\r\n")); err != nil {