diff --git a/proxy/udp.go b/proxy/udp.go index 543962d..4f47466 100644 --- a/proxy/udp.go +++ b/proxy/udp.go @@ -20,8 +20,7 @@ type udpHandler struct { proxyPort int timeout time.Duration - remoteAddrMap sync.Map - remoteConnMap sync.Map + remoteMap sync.Map } func NewUDPHandler(proxyHost string, proxyPort int, timeout time.Duration) core.UDPConnHandler { @@ -96,8 +95,10 @@ func (h *udpHandler) Connect(conn core.UDPConn, target *net.UDPAddr) error { remoteConn = &S.PacketConn{Session: session, PacketConn: remoteConn} } - h.remoteAddrMap.Store(conn, remoteAddr) - h.remoteConnMap.Store(conn, remoteConn) + h.remoteMap.Store(conn, &udpElement{ + remoteAddr: remoteAddr, + remoteConn: remoteConn, + }) go h.fetchUDPInput(conn, remoteConn, target) @@ -130,12 +131,9 @@ func (h *udpHandler) ReceiveTo(conn core.UDPConn, data []byte, addr *net.UDPAddr var remoteAddr net.Addr var remoteConn net.PacketConn - if value, ok := h.remoteAddrMap.Load(conn); ok { - remoteAddr = value.(net.Addr) - } - - if value, ok := h.remoteConnMap.Load(conn); ok { - remoteConn = value.(net.PacketConn) + if elm, ok := h.remoteMap.Load(conn); ok { + remoteAddr = elm.(*udpElement).remoteAddr + remoteConn = elm.(*udpElement).remoteConn } if remoteAddr == nil || remoteConn == nil { @@ -154,15 +152,13 @@ func (h *udpHandler) Close(conn core.UDPConn) { conn.Close() // Load from remoteConnMap - if remoteConn, ok := h.remoteConnMap.Load(conn); ok { - remoteConn.(net.PacketConn).Close() + if elm, ok := h.remoteMap.Load(conn); ok { + elm.(*udpElement).remoteConn.Close() + h.remoteMap.Delete(conn) } else { return } - h.remoteAddrMap.Delete(conn) - h.remoteConnMap.Delete(conn) - // Remove session removeSession(conn) } diff --git a/proxy/utils.go b/proxy/utils.go index a97d33b..a2d64be 100644 --- a/proxy/utils.go +++ b/proxy/utils.go @@ -10,6 +10,12 @@ import ( "github.com/xjasonlyu/tun2socks/proxy/socks" ) +// UDP util +type udpElement struct { + remoteAddr net.Addr + remoteConn net.PacketConn +} + // TCP functions type duplexConn interface { net.Conn