From cc4310106b61ba9d21b11629d7e5a18eb45e832e Mon Sep 17 00:00:00 2001 From: ginuerzh Date: Thu, 19 Oct 2023 23:47:47 +0800 Subject: [PATCH] fix race condition --- connector/relay/conn.go | 8 ++--- connector/relay/metadata.go | 7 ----- connector/socks/v5/bind.go | 2 +- connector/socks/v5/conn.go | 12 +++---- connector/socks/v5/metadata.go | 11 +++++++ connector/ss/connector.go | 6 ++-- connector/tunnel/conn.go | 8 ++--- connector/tunnel/metadata.go | 7 ----- handler/dns/handler.go | 16 +++++----- handler/socks/v5/mbind.go | 2 +- handler/socks/v5/metadata.go | 12 +++++++ handler/ss/udp/handler.go | 8 ++--- handler/tap/handler.go | 30 +++++++++--------- handler/tun/client.go | 42 ++++++++++++------------- handler/tun/server.go | 36 ++++++++++----------- handler/tunnel/metadata.go | 7 ----- internal/net/transport.go | 2 +- internal/net/udp/relay.go | 8 ++--- internal/util/dtls/conn.go | 6 ++-- internal/util/icmp/conn.go | 16 +++++----- internal/util/mux/mux.go | 6 ++++ internal/util/pht/server.go | 8 ++--- internal/util/relay/conn.go | 8 ++--- internal/util/socks/conn.go | 8 ++--- internal/util/ss/conn.go | 12 +++---- listener/mws/metadata.go | 7 ----- listener/redirect/udp/conn.go | 2 +- listener/redirect/udp/listener_linux.go | 8 ++--- listener/tun/tun.go | 8 ++--- 29 files changed, 157 insertions(+), 156 deletions(-) diff --git a/connector/relay/conn.go b/connector/relay/conn.go index 2ce2263..1712051 100644 --- a/connector/relay/conn.go +++ b/connector/relay/conn.go @@ -75,8 +75,8 @@ func (c *udpConn) Read(b []byte) (n int, err error) { buf := bufpool.Get(dlen) defer bufpool.Put(buf) - _, err = io.ReadFull(c.Conn, *buf) - n = copy(b, *buf) + _, err = io.ReadFull(c.Conn, buf) + n = copy(b, buf) return } @@ -169,8 +169,8 @@ func (c *bindUDPConn) Read(b []byte) (n int, err error) { buf := bufpool.Get(dlen) defer bufpool.Put(buf) - _, err = io.ReadFull(c.Conn, *buf) - n = copy(b, *buf) + _, err = io.ReadFull(c.Conn, buf) + n = copy(b, buf) return } diff --git a/connector/relay/metadata.go b/connector/relay/metadata.go index 905adae..8cf22af 100644 --- a/connector/relay/metadata.go +++ b/connector/relay/metadata.go @@ -10,10 +10,6 @@ import ( "github.com/google/uuid" ) -const ( - defaultMuxVersion = 1 -) - type metadata struct { connectTimeout time.Duration noDelay bool @@ -47,9 +43,6 @@ func (c *relayConnector) parseMetadata(md mdata.Metadata) (err error) { MaxReceiveBuffer: mdutil.GetInt(md, "mux.maxReceiveBuffer"), MaxStreamBuffer: mdutil.GetInt(md, "mux.maxStreamBuffer"), } - if c.md.muxCfg.Version == 0 { - c.md.muxCfg.Version = defaultMuxVersion - } return } diff --git a/connector/socks/v5/bind.go b/connector/socks/v5/bind.go index 6922880..31f0ea8 100644 --- a/connector/socks/v5/bind.go +++ b/connector/socks/v5/bind.go @@ -62,7 +62,7 @@ func (c *socks5Connector) muxBindTCP(ctx context.Context, conn net.Conn, network return nil, err } - session, err := mux.ServerSession(conn, nil) + session, err := mux.ServerSession(conn, c.md.muxCfg) if err != nil { return nil, err } diff --git a/connector/socks/v5/conn.go b/connector/socks/v5/conn.go index 3735ca8..da0552f 100644 --- a/connector/socks/v5/conn.go +++ b/connector/socks/v5/conn.go @@ -36,7 +36,7 @@ func (c *udpRelayConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { buf := bufpool.Get(c.bufferSize) defer bufpool.Put(buf) - nn, err := c.udpConn.Read(*buf) + nn, err := c.udpConn.Read(buf) if err != nil { return } @@ -48,7 +48,7 @@ func (c *udpRelayConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { dgram := gosocks5.UDPDatagram{ Header: &header, } - _, err = dgram.ReadFrom(bytes.NewReader((*buf)[:nn])) + _, err = dgram.ReadFrom(bytes.NewReader(buf[:nn])) if err != nil { return } @@ -81,15 +81,15 @@ func (c *udpRelayConn) WriteTo(b []byte, addr net.Addr) (n int, err error) { buf := bufpool.Get(c.bufferSize) defer bufpool.Put(buf) - nn, err := dgram.WriteTo(bytes.NewBuffer((*buf)[:0])) + nn, err := dgram.WriteTo(bytes.NewBuffer(buf[:0])) if err != nil { return } - if nn > int64(len(*buf)) { - nn = int64(len(*buf)) + if nn > int64(len(buf)) { + nn = int64(len(buf)) } - _, err = c.udpConn.Write((*buf)[:nn]) + _, err = c.udpConn.Write(buf[:nn]) n = len(b) return diff --git a/connector/socks/v5/metadata.go b/connector/socks/v5/metadata.go index 8428e14..0ee0616 100644 --- a/connector/socks/v5/metadata.go +++ b/connector/socks/v5/metadata.go @@ -5,6 +5,7 @@ import ( mdata "github.com/go-gost/core/metadata" mdutil "github.com/go-gost/core/metadata/util" + "github.com/go-gost/x/internal/util/mux" ) const ( @@ -16,6 +17,7 @@ type metadata struct { noTLS bool relay string udpBufferSize int + muxCfg *mux.Config } func (c *socks5Connector) parseMetadata(md mdata.Metadata) (err error) { @@ -34,5 +36,14 @@ func (c *socks5Connector) parseMetadata(md mdata.Metadata) (err error) { c.md.udpBufferSize = defaultUDPBufferSize } + c.md.muxCfg = &mux.Config{ + Version: mdutil.GetInt(md, "mux.version"), + KeepAliveInterval: mdutil.GetDuration(md, "mux.keepaliveInterval"), + KeepAliveDisabled: mdutil.GetBool(md, "mux.keepaliveDisabled"), + KeepAliveTimeout: mdutil.GetDuration(md, "mux.keepaliveTimeout"), + MaxFrameSize: mdutil.GetInt(md, "mux.maxFrameSize"), + MaxReceiveBuffer: mdutil.GetInt(md, "mux.maxReceiveBuffer"), + MaxStreamBuffer: mdutil.GetInt(md, "mux.maxStreamBuffer"), + } return } diff --git a/connector/ss/connector.go b/connector/ss/connector.go index b6f422d..bc36397 100644 --- a/connector/ss/connector.go +++ b/connector/ss/connector.go @@ -80,7 +80,7 @@ func (c *ssConnector) Connect(ctx context.Context, conn net.Conn, network, addre rawaddr := bufpool.Get(512) defer bufpool.Put(rawaddr) - n, err := addr.Encode(*rawaddr) + n, err := addr.Encode(rawaddr) if err != nil { log.Error("encoding addr: ", err) return nil, err @@ -99,12 +99,12 @@ func (c *ssConnector) Connect(ctx context.Context, conn net.Conn, network, addre if c.md.noDelay { sc = ss.ShadowConn(conn, nil) // write the addr at once. - if _, err := sc.Write((*rawaddr)[:n]); err != nil { + if _, err := sc.Write(rawaddr[:n]); err != nil { return nil, err } } else { // cache the header - sc = ss.ShadowConn(conn, (*rawaddr)[:n]) + sc = ss.ShadowConn(conn, rawaddr[:n]) } return sc, nil diff --git a/connector/tunnel/conn.go b/connector/tunnel/conn.go index 59cd0c1..cd712b2 100644 --- a/connector/tunnel/conn.go +++ b/connector/tunnel/conn.go @@ -75,8 +75,8 @@ func (c *udpConn) Read(b []byte) (n int, err error) { buf := bufpool.Get(dlen) defer bufpool.Put(buf) - _, err = io.ReadFull(c.Conn, *buf) - n = copy(b, *buf) + _, err = io.ReadFull(c.Conn, buf) + n = copy(b, buf) return } @@ -169,8 +169,8 @@ func (c *bindUDPConn) Read(b []byte) (n int, err error) { buf := bufpool.Get(dlen) defer bufpool.Put(buf) - _, err = io.ReadFull(c.Conn, *buf) - n = copy(b, *buf) + _, err = io.ReadFull(c.Conn, buf) + n = copy(b, buf) return } diff --git a/connector/tunnel/metadata.go b/connector/tunnel/metadata.go index dee4472..4257a6b 100644 --- a/connector/tunnel/metadata.go +++ b/connector/tunnel/metadata.go @@ -11,10 +11,6 @@ import ( "github.com/google/uuid" ) -const ( - defaultMuxVersion = 2 -) - var ( ErrInvalidTunnelID = errors.New("tunnel: invalid tunnel ID") ) @@ -51,9 +47,6 @@ func (c *tunnelConnector) parseMetadata(md mdata.Metadata) (err error) { MaxReceiveBuffer: mdutil.GetInt(md, "mux.maxReceiveBuffer"), MaxStreamBuffer: mdutil.GetInt(md, "mux.maxStreamBuffer"), } - if c.md.muxCfg.Version == 0 { - c.md.muxCfg.Version = defaultMuxVersion - } return } diff --git a/handler/dns/handler.go b/handler/dns/handler.go index 7cb004b..18ec098 100644 --- a/handler/dns/handler.go +++ b/handler/dns/handler.go @@ -141,17 +141,17 @@ func (h *dnsHandler) Handle(ctx context.Context, conn net.Conn, opts ...handler. b := bufpool.Get(h.md.bufferSize) defer bufpool.Put(b) - n, err := conn.Read(*b) + n, err := conn.Read(b) if err != nil { log.Error(err) return err } - reply, err := h.request(ctx, (*b)[:n], log) + reply, err := h.request(ctx, b[:n], log) if err != nil { return err } - defer bufpool.Put(&reply) + defer bufpool.Put(reply) if _, err = conn.Write(reply); err != nil { log.Error(err) @@ -203,14 +203,14 @@ func (h *dnsHandler) request(ctx context.Context, msg []byte, log logger.Logger) log.Debug("bypass: ", mq.Question[0].Name) mr = (&dns.Msg{}).SetReply(&mq) b := bufpool.Get(h.md.bufferSize) - return mr.PackBuffer(*b) + return mr.PackBuffer(b) } } mr = h.lookupHosts(ctx, &mq, log) if mr != nil { b := bufpool.Get(h.md.bufferSize) - return mr.PackBuffer(*b) + return mr.PackBuffer(b) } // only cache for single question message. @@ -222,14 +222,14 @@ func (h *dnsHandler) request(ctx context.Context, msg []byte, log logger.Logger) if int32(ttl.Seconds()) > 0 { log.Debugf("message %d (cached): %s", mq.Id, mq.Question[0].String()) b := bufpool.Get(h.md.bufferSize) - return mr.PackBuffer(*b) + return mr.PackBuffer(b) } } } if mr != nil && h.md.async { b := bufpool.Get(h.md.bufferSize) - reply, err := mr.PackBuffer(*b) + reply, err := mr.PackBuffer(b) if err != nil { return nil, err } @@ -248,7 +248,7 @@ func (h *dnsHandler) exchange(ctx context.Context, mq *dns.Msg) ([]byte, error) b := bufpool.Get(h.md.bufferSize) defer bufpool.Put(b) - query, err := mq.PackBuffer(*b) + query, err := mq.PackBuffer(b) if err != nil { return nil, err } diff --git a/handler/socks/v5/mbind.go b/handler/socks/v5/mbind.go index a5ac3d3..9038780 100644 --- a/handler/socks/v5/mbind.go +++ b/handler/socks/v5/mbind.go @@ -70,7 +70,7 @@ func (h *socks5Handler) muxBindLocal(ctx context.Context, conn net.Conn, network func (h *socks5Handler) serveMuxBind(ctx context.Context, conn net.Conn, ln net.Listener, log logger.Logger) error { // Upgrade connection to multiplex stream. - session, err := mux.ClientSession(conn, nil) + session, err := mux.ClientSession(conn, h.md.muxCfg) if err != nil { log.Error(err) return err diff --git a/handler/socks/v5/metadata.go b/handler/socks/v5/metadata.go index 27ca67e..ec31575 100644 --- a/handler/socks/v5/metadata.go +++ b/handler/socks/v5/metadata.go @@ -6,6 +6,7 @@ import ( mdata "github.com/go-gost/core/metadata" mdutil "github.com/go-gost/core/metadata/util" + "github.com/go-gost/x/internal/util/mux" ) type metadata struct { @@ -16,6 +17,7 @@ type metadata struct { udpBufferSize int compatibilityMode bool hash string + muxCfg *mux.Config } func (h *socks5Handler) parseMetadata(md mdata.Metadata) (err error) { @@ -43,5 +45,15 @@ func (h *socks5Handler) parseMetadata(md mdata.Metadata) (err error) { h.md.compatibilityMode = mdutil.GetBool(md, compatibilityMode) h.md.hash = mdutil.GetString(md, hash) + h.md.muxCfg = &mux.Config{ + Version: mdutil.GetInt(md, "mux.version"), + KeepAliveInterval: mdutil.GetDuration(md, "mux.keepaliveInterval"), + KeepAliveDisabled: mdutil.GetBool(md, "mux.keepaliveDisabled"), + KeepAliveTimeout: mdutil.GetDuration(md, "mux.keepaliveTimeout"), + MaxFrameSize: mdutil.GetInt(md, "mux.maxFrameSize"), + MaxReceiveBuffer: mdutil.GetInt(md, "mux.maxReceiveBuffer"), + MaxStreamBuffer: mdutil.GetInt(md, "mux.maxStreamBuffer"), + } + return nil } diff --git a/handler/ss/udp/handler.go b/handler/ss/udp/handler.go index 2143709..4f31422 100644 --- a/handler/ss/udp/handler.go +++ b/handler/ss/udp/handler.go @@ -130,7 +130,7 @@ func (h *ssuHandler) relayPacket(pc1, pc2 net.PacketConn, log logger.Logger) (er b := bufpool.Get(bufSize) defer bufpool.Put(b) - n, addr, err := pc1.ReadFrom(*b) + n, addr, err := pc1.ReadFrom(b) if err != nil { return err } @@ -140,7 +140,7 @@ func (h *ssuHandler) relayPacket(pc1, pc2 net.PacketConn, log logger.Logger) (er return nil } - if _, err = pc2.WriteTo((*b)[:n], addr); err != nil { + if _, err = pc2.WriteTo(b[:n], addr); err != nil { return err } @@ -162,7 +162,7 @@ func (h *ssuHandler) relayPacket(pc1, pc2 net.PacketConn, log logger.Logger) (er b := bufpool.Get(bufSize) defer bufpool.Put(b) - n, raddr, err := pc2.ReadFrom(*b) + n, raddr, err := pc2.ReadFrom(b) if err != nil { return err } @@ -172,7 +172,7 @@ func (h *ssuHandler) relayPacket(pc1, pc2 net.PacketConn, log logger.Logger) (er return nil } - if _, err = pc1.WriteTo((*b)[:n], raddr); err != nil { + if _, err = pc1.WriteTo(b[:n], raddr); err != nil { return err } diff --git a/handler/tap/handler.go b/handler/tap/handler.go index 1d081a2..ccd1dc2 100644 --- a/handler/tap/handler.go +++ b/handler/tap/handler.go @@ -196,7 +196,7 @@ func (h *tapHandler) transport(tap net.Conn, conn net.PacketConn, raddr net.Addr b := bufpool.Get(h.md.bufferSize) defer bufpool.Put(b) - n, err := tap.Read(*b) + n, err := tap.Read(b) if err != nil { select { case h.exit <- struct{}{}: @@ -208,22 +208,22 @@ func (h *tapHandler) transport(tap net.Conn, conn net.PacketConn, raddr net.Addr return nil } - src := waterutil.MACSource((*b)[:n]) - dst := waterutil.MACDestination((*b)[:n]) - eType := etherType(waterutil.MACEthertype((*b)[:n])) + src := waterutil.MACSource(b[:n]) + dst := waterutil.MACDestination(b[:n]) + eType := etherType(waterutil.MACEthertype(b[:n])) log.Debugf("%s >> %s %s %d", src, dst, eType, n) // client side, deliver frame directly. if raddr != nil { - _, err := conn.WriteTo((*b)[:n], raddr) + _, err := conn.WriteTo(b[:n], raddr) return err } // server side, broadcast. if waterutil.IsBroadcast(dst) { go h.routes.Range(func(k, v any) bool { - conn.WriteTo((*b)[:n], v.(net.Addr)) + conn.WriteTo(b[:n], v.(net.Addr)) return true }) return nil @@ -238,7 +238,7 @@ func (h *tapHandler) transport(tap net.Conn, conn net.PacketConn, raddr net.Addr return nil } - if _, err := conn.WriteTo((*b)[:n], addr); err != nil { + if _, err := conn.WriteTo(b[:n], addr); err != nil { return err } @@ -258,7 +258,7 @@ func (h *tapHandler) transport(tap net.Conn, conn net.PacketConn, raddr net.Addr b := bufpool.Get(h.md.bufferSize) defer bufpool.Put(b) - n, addr, err := conn.ReadFrom(*b) + n, addr, err := conn.ReadFrom(b) if err != nil && err != shadowaead.ErrShortPacket { return err @@ -267,15 +267,15 @@ func (h *tapHandler) transport(tap net.Conn, conn net.PacketConn, raddr net.Addr return nil } - src := waterutil.MACSource((*b)[:n]) - dst := waterutil.MACDestination((*b)[:n]) - eType := etherType(waterutil.MACEthertype((*b)[:n])) + src := waterutil.MACSource(b[:n]) + dst := waterutil.MACDestination(b[:n]) + eType := etherType(waterutil.MACEthertype(b[:n])) log.Debugf("%s >> %s %s %d", src, dst, eType, n) // client side, deliver frame to tap device. if raddr != nil { - _, err := tap.Write((*b)[:n]) + _, err := tap.Write(b[:n]) return err } @@ -294,7 +294,7 @@ func (h *tapHandler) transport(tap net.Conn, conn net.PacketConn, raddr net.Addr if waterutil.IsBroadcast(dst) { go h.routes.Range(func(k, v any) bool { if k.(tapRouteKey) != rkey { - conn.WriteTo((*b)[:n], v.(net.Addr)) + conn.WriteTo(b[:n], v.(net.Addr)) } return true }) @@ -302,11 +302,11 @@ func (h *tapHandler) transport(tap net.Conn, conn net.PacketConn, raddr net.Addr if v, ok := h.routes.Load(hwAddrToTapRouteKey(dst)); ok { log.Debugf("find route: %s -> %s", dst, v) - _, err := conn.WriteTo((*b)[:n], v.(net.Addr)) + _, err := conn.WriteTo(b[:n], v.(net.Addr)) return err } - if _, err := tap.Write((*b)[:n]); err != nil { + if _, err := tap.Write(b[:n]); err != nil { select { case h.exit <- struct{}{}: default: diff --git a/handler/tun/client.go b/handler/tun/client.go index 3eab336..af3e4b0 100644 --- a/handler/tun/client.go +++ b/handler/tun/client.go @@ -62,14 +62,14 @@ func (h *tunHandler) keepAlive(ctx context.Context, conn net.Conn, ips []net.IP) keepAliveData := bufpool.Get(keepAliveHeaderLength + len(ips)*net.IPv6len) defer bufpool.Put(keepAliveData) - copy((*keepAliveData)[:4], magicHeader) // magic header - copy((*keepAliveData)[4:20], []byte(h.md.passphrase)) + copy(keepAliveData[:4], magicHeader) // magic header + copy(keepAliveData[4:20], []byte(h.md.passphrase)) pos := 20 for _, ip := range ips { - copy((*keepAliveData)[pos:pos+net.IPv6len], ip.To16()) + copy(keepAliveData[pos:pos+net.IPv6len], ip.To16()) pos += net.IPv6len } - if _, err := conn.Write((*keepAliveData)); err != nil { + if _, err := conn.Write(keepAliveData); err != nil { return } @@ -84,7 +84,7 @@ func (h *tunHandler) keepAlive(ctx context.Context, conn net.Conn, ips []net.IP) for { select { case <-ticker.C: - if _, err := conn.Write((*keepAliveData)); err != nil { + if _, err := conn.Write(keepAliveData); err != nil { return } h.options.Logger.Debugf("keepalive sended") @@ -103,23 +103,23 @@ func (h *tunHandler) transportClient(tun io.ReadWriter, conn net.Conn, log logge b := bufpool.Get(h.md.bufferSize) defer bufpool.Put(b) - n, err := tun.Read(*b) + n, err := tun.Read(b) if err != nil { return ErrTun } - if waterutil.IsIPv4((*b)[:n]) { - header, err := ipv4.ParseHeader((*b)[:n]) + if waterutil.IsIPv4(b[:n]) { + header, err := ipv4.ParseHeader(b[:n]) if err != nil { log.Warn(err) return nil } log.Tracef("%s >> %s %-4s %d/%-4d %-4x %d", - header.Src, header.Dst, ipProtocol(waterutil.IPv4Protocol((*b)[:n])), + header.Src, header.Dst, ipProtocol(waterutil.IPv4Protocol(b[:n])), header.Len, header.TotalLen, header.ID, header.Flags) - } else if waterutil.IsIPv6((*b)[:n]) { - header, err := ipv6.ParseHeader((*b)[:n]) + } else if waterutil.IsIPv6(b[:n]) { + header, err := ipv6.ParseHeader(b[:n]) if err != nil { log.Warn(err) return nil @@ -134,7 +134,7 @@ func (h *tunHandler) transportClient(tun io.ReadWriter, conn net.Conn, log logge return nil } - _, err = conn.Write((*b)[:n]) + _, err = conn.Write(b[:n]) return err }() @@ -151,13 +151,13 @@ func (h *tunHandler) transportClient(tun io.ReadWriter, conn net.Conn, log logge b := bufpool.Get(h.md.bufferSize) defer bufpool.Put(b) - n, err := conn.Read(*b) + n, err := conn.Read(b) if err != nil { return err } - if n == keepAliveHeaderLength && bytes.Equal((*b)[:4], magicHeader) { - ip := net.IP((*b)[4:20]) + if n == keepAliveHeaderLength && bytes.Equal(b[:4], magicHeader) { + ip := net.IP(b[4:20]) log.Debugf("keepalive received at %v", ip) if h.md.keepAlivePeriod > 0 { @@ -166,18 +166,18 @@ func (h *tunHandler) transportClient(tun io.ReadWriter, conn net.Conn, log logge return nil } - if waterutil.IsIPv4((*b)[:n]) { - header, err := ipv4.ParseHeader((*b)[:n]) + if waterutil.IsIPv4(b[:n]) { + header, err := ipv4.ParseHeader(b[:n]) if err != nil { log.Warn(err) return nil } log.Tracef("%s >> %s %-4s %d/%-4d %-4x %d", - header.Src, header.Dst, ipProtocol(waterutil.IPv4Protocol((*b)[:n])), + header.Src, header.Dst, ipProtocol(waterutil.IPv4Protocol(b[:n])), header.Len, header.TotalLen, header.ID, header.Flags) - } else if waterutil.IsIPv6((*b)[:n]) { - header, err := ipv6.ParseHeader((*b)[:n]) + } else if waterutil.IsIPv6(b[:n]) { + header, err := ipv6.ParseHeader(b[:n]) if err != nil { log.Warn(err) return nil @@ -192,7 +192,7 @@ func (h *tunHandler) transportClient(tun io.ReadWriter, conn net.Conn, log logge return nil } - if _, err = tun.Write((*b)[:n]); err != nil { + if _, err = tun.Write(b[:n]); err != nil { return ErrTun } return nil diff --git a/handler/tun/server.go b/handler/tun/server.go index 09e7b66..22a13d6 100644 --- a/handler/tun/server.go +++ b/handler/tun/server.go @@ -45,7 +45,7 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con b := bufpool.Get(h.md.bufferSize) defer bufpool.Put(b) - n, err := tun.Read(*b) + n, err := tun.Read(b) if err != nil { return ErrTun } @@ -54,8 +54,8 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con } var src, dst net.IP - if waterutil.IsIPv4((*b)[:n]) { - header, err := ipv4.ParseHeader((*b)[:n]) + if waterutil.IsIPv4(b[:n]) { + header, err := ipv4.ParseHeader(b[:n]) if err != nil { log.Warnf("parse ipv4 packet header: %v", err) return nil @@ -63,10 +63,10 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con src, dst = header.Src, header.Dst log.Tracef("%s >> %s %-4s %d/%-4d %-4x %d", - src, dst, ipProtocol(waterutil.IPv4Protocol((*b)[:n])), + src, dst, ipProtocol(waterutil.IPv4Protocol(b[:n])), header.Len, header.TotalLen, header.ID, header.Flags) - } else if waterutil.IsIPv6((*b)[:n]) { - header, err := ipv6.ParseHeader((*b)[:n]) + } else if waterutil.IsIPv6(b[:n]) { + header, err := ipv6.ParseHeader(b[:n]) if err != nil { log.Warnf("parse ipv6 packet header: %v", err) return nil @@ -90,7 +90,7 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con log.Debugf("find route: %s -> %s", dst, addr) - if _, err := conn.WriteTo((*b)[:n], addr); err != nil { + if _, err := conn.WriteTo(b[:n], addr); err != nil { return err } return nil @@ -109,16 +109,16 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con b := bufpool.Get(h.md.bufferSize) defer bufpool.Put(b) - n, addr, err := conn.ReadFrom(*b) + n, addr, err := conn.ReadFrom(b) if err != nil { return err } if n == 0 { return nil } - if n > keepAliveHeaderLength && bytes.Equal((*b)[:4], magicHeader) { + if n > keepAliveHeaderLength && bytes.Equal(b[:4], magicHeader) { var peerIPs []net.IP - data := (*b)[keepAliveHeaderLength:n] + data := b[keepAliveHeaderLength:n] if len(data)%net.IPv6len == 0 { for len(data) > 0 { peerIPs = append(peerIPs, net.IP(data[:net.IPv6len])) @@ -139,7 +139,7 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con if auther := h.options.Auther; auther != nil { ok := true - key := bytes.TrimRight((*b)[4:20], "\x00") + key := bytes.TrimRight(b[4:20], "\x00") for _, ip := range peerIPs { if _, ok = auther.Authenticate(ctx, ip.String(), string(key)); !ok { break @@ -175,8 +175,8 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con } var src, dst net.IP - if waterutil.IsIPv4((*b)[:n]) { - header, err := ipv4.ParseHeader((*b)[:n]) + if waterutil.IsIPv4(b[:n]) { + header, err := ipv4.ParseHeader(b[:n]) if err != nil { log.Warnf("parse ipv4 packet header: %v", err) return nil @@ -184,10 +184,10 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con src, dst = header.Src, header.Dst log.Tracef("%s >> %s %-4s %d/%-4d %-4x %d", - src, dst, ipProtocol(waterutil.IPv4Protocol((*b)[:n])), + src, dst, ipProtocol(waterutil.IPv4Protocol(b[:n])), header.Len, header.TotalLen, header.ID, header.Flags) - } else if waterutil.IsIPv6((*b)[:n]) { - header, err := ipv6.ParseHeader((*b)[:n]) + } else if waterutil.IsIPv6(b[:n]) { + header, err := ipv6.ParseHeader(b[:n]) if err != nil { log.Warnf("parse ipv6 packet header: %v", err) return nil @@ -206,11 +206,11 @@ func (h *tunHandler) transportServer(ctx context.Context, tun io.ReadWriter, con if addr := h.findRouteFor(dst, config.Routes...); addr != nil { log.Debugf("find route: %s -> %s", dst, addr) - _, err := conn.WriteTo((*b)[:n], addr) + _, err := conn.WriteTo(b[:n], addr) return err } - if _, err := tun.Write((*b)[:n]); err != nil { + if _, err := tun.Write(b[:n]); err != nil { return ErrTun } return nil diff --git a/handler/tunnel/metadata.go b/handler/tunnel/metadata.go index e94b395..85a0bd2 100644 --- a/handler/tunnel/metadata.go +++ b/handler/tunnel/metadata.go @@ -14,10 +14,6 @@ import ( "github.com/go-gost/x/registry" ) -const ( - defaultMuxVersion = 2 -) - type metadata struct { readTimeout time.Duration noDelay bool @@ -69,9 +65,6 @@ func (h *tunnelHandler) parseMetadata(md mdata.Metadata) (err error) { MaxReceiveBuffer: mdutil.GetInt(md, "mux.maxReceiveBuffer"), MaxStreamBuffer: mdutil.GetInt(md, "mux.maxStreamBuffer"), } - if h.md.muxCfg.Version == 0 { - h.md.muxCfg.Version = defaultMuxVersion - } return } diff --git a/internal/net/transport.go b/internal/net/transport.go index ee4b5d1..7ce4b10 100644 --- a/internal/net/transport.go +++ b/internal/net/transport.go @@ -33,7 +33,7 @@ func CopyBuffer(dst io.Writer, src io.Reader, bufSize int) error { buf := bufpool.Get(bufSize) defer bufpool.Put(buf) - _, err := io.CopyBuffer(dst, src, *buf) + _, err := io.CopyBuffer(dst, src, buf) return err } diff --git a/internal/net/udp/relay.go b/internal/net/udp/relay.go index ddf3171..79146f1 100644 --- a/internal/net/udp/relay.go +++ b/internal/net/udp/relay.go @@ -53,7 +53,7 @@ func (r *Relay) Run() (err error) { b := bufpool.Get(bufSize) defer bufpool.Put(b) - n, raddr, err := r.pc1.ReadFrom(*b) + n, raddr, err := r.pc1.ReadFrom(b) if err != nil { return err } @@ -65,7 +65,7 @@ func (r *Relay) Run() (err error) { return nil } - if _, err := r.pc2.WriteTo((*b)[:n], raddr); err != nil { + if _, err := r.pc2.WriteTo(b[:n], raddr); err != nil { return err } @@ -91,7 +91,7 @@ func (r *Relay) Run() (err error) { b := bufpool.Get(bufSize) defer bufpool.Put(b) - n, raddr, err := r.pc2.ReadFrom(*b) + n, raddr, err := r.pc2.ReadFrom(b) if err != nil { return err } @@ -103,7 +103,7 @@ func (r *Relay) Run() (err error) { return nil } - if _, err := r.pc1.WriteTo((*b)[:n], raddr); err != nil { + if _, err := r.pc1.WriteTo(b[:n], raddr); err != nil { return err } diff --git a/internal/util/dtls/conn.go b/internal/util/dtls/conn.go index c120135..e5b5fea 100644 --- a/internal/util/dtls/conn.go +++ b/internal/util/dtls/conn.go @@ -39,13 +39,13 @@ func (c *dtlsConn) Read(p []byte) (n int, err error) { buf := bufpool.Get(bufferSize) defer bufpool.Put(buf) - nn, err := c.Conn.Read(*buf) + nn, err := c.Conn.Read(buf) if err != nil { return 0, err } - n = copy(p, (*buf)[:nn]) - c.rbuf.Write((*buf)[n:nn]) + n = copy(p, buf[:nn]) + c.rbuf.Write(buf[n:nn]) return } diff --git a/internal/util/icmp/conn.go b/internal/util/icmp/conn.go index 5d333aa..593e421 100644 --- a/internal/util/icmp/conn.go +++ b/internal/util/icmp/conn.go @@ -96,12 +96,12 @@ func (c *clientConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { defer bufpool.Put(buf) for { - n, addr, err = c.PacketConn.ReadFrom(*buf) + n, addr, err = c.PacketConn.ReadFrom(buf) if err != nil { return } - m, err := icmp.ParseMessage(1, (*buf)[:n]) + m, err := icmp.ParseMessage(1, buf[:n]) if err != nil { // logger.Default().Error("icmp: parse message %v", err) return 0, addr, err @@ -155,7 +155,7 @@ func (c *clientConn) WriteTo(b []byte, addr net.Addr) (n int, err error) { msg := message{ data: b, } - nn, err := msg.Encode(*buf) + nn, err := msg.Encode(buf) if err != nil { return } @@ -163,7 +163,7 @@ func (c *clientConn) WriteTo(b []byte, addr net.Addr) (n int, err error) { echo := icmp.Echo{ ID: c.id, Seq: int(atomic.AddUint32(&c.seq, 1)), - Data: (*buf)[:nn], + Data: buf[:nn], } m := icmp.Message{ Type: ipv4.ICMPTypeEcho, @@ -195,12 +195,12 @@ func (c *serverConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { defer bufpool.Put(buf) for { - n, addr, err = c.PacketConn.ReadFrom(*buf) + n, addr, err = c.PacketConn.ReadFrom(buf) if err != nil { return } - m, err := icmp.ParseMessage(1, (*buf)[:n]) + m, err := icmp.ParseMessage(1, buf[:n]) if err != nil { // logger.Default().Error("icmp: parse message %v", err) return 0, addr, err @@ -260,7 +260,7 @@ func (c *serverConn) WriteTo(b []byte, addr net.Addr) (n int, err error) { flags: FlagAck, data: b, } - nn, err := msg.Encode(*buf) + nn, err := msg.Encode(buf) if err != nil { return } @@ -268,7 +268,7 @@ func (c *serverConn) WriteTo(b []byte, addr net.Addr) (n int, err error) { echo := icmp.Echo{ ID: id, Seq: int(atomic.LoadUint32(&c.seqs[id-1])), - Data: (*buf)[:nn], + Data: buf[:nn], } m := icmp.Message{ Type: ipv4.ICMPTypeEchoReply, diff --git a/internal/util/mux/mux.go b/internal/util/mux/mux.go index 3aaf483..15be88e 100644 --- a/internal/util/mux/mux.go +++ b/internal/util/mux/mux.go @@ -7,6 +7,10 @@ import ( smux "github.com/xtaci/smux" ) +const ( + defaultVersion = 2 +) + type Config struct { // SMUX Protocol version, support 1,2 Version int @@ -36,6 +40,8 @@ type Config struct { func convertConfig(cfg *Config) *smux.Config { smuxCfg := smux.DefaultConfig() + smuxCfg.Version = defaultVersion + if cfg == nil { return smuxCfg } diff --git a/internal/util/pht/server.go b/internal/util/pht/server.go index 1c28038..411e4aa 100644 --- a/internal/util/pht/server.go +++ b/internal/util/pht/server.go @@ -375,10 +375,10 @@ func (s *Server) handlePull(w http.ResponseWriter, r *http.Request) { for { conn.SetReadDeadline(time.Now().Add(s.options.readTimeout)) - n, err := conn.Read(*b) + n, err := conn.Read(b) if n > 0 { bw := bufio.NewWriter(w) - bw.WriteString(base64.StdEncoding.EncodeToString((*b)[:n])) + bw.WriteString(base64.StdEncoding.EncodeToString(b[:n])) bw.WriteString("\n") if err := bw.Flush(); err != nil { return @@ -389,8 +389,8 @@ func (s *Server) handlePull(w http.ResponseWriter, r *http.Request) { } if err != nil { if errors.Is(err, os.ErrDeadlineExceeded) { - (*b)[0] = '\n' // no data - w.Write((*b)[:1]) + b[0] = '\n' // no data + w.Write(b[:1]) } else if errors.Is(err, io.EOF) { // server connection closed } else { diff --git a/internal/util/relay/conn.go b/internal/util/relay/conn.go index 8d4c671..2d25125 100644 --- a/internal/util/relay/conn.go +++ b/internal/util/relay/conn.go @@ -110,7 +110,7 @@ func (c *udpConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { rbuf := bufpool.Get(c.bufferSize) defer bufpool.Put(rbuf) - n, c.raddr, err = c.PacketConn.ReadFrom(*rbuf) + n, c.raddr, err = c.PacketConn.ReadFrom(rbuf) if err != nil { return } @@ -119,11 +119,11 @@ func (c *udpConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { header := gosocks5.UDPHeader{ Addr: &socksAddr, } - hlen, err := header.ReadFrom(bytes.NewReader((*rbuf)[:n])) + hlen, err := header.ReadFrom(bytes.NewReader(rbuf[:n])) if err != nil { return } - n = copy(b, (*rbuf)[hlen:n]) + n = copy(b, rbuf[hlen:n]) addr, err = net.ResolveUDPAddr("udp", socksAddr.String()) return @@ -151,7 +151,7 @@ func (c *udpConn) WriteTo(b []byte, addr net.Addr) (n int, err error) { Data: b, } - buf := bytes.NewBuffer((*wbuf)[:0]) + buf := bytes.NewBuffer(wbuf[:0]) _, err = dgram.WriteTo(buf) if err != nil { return diff --git a/internal/util/socks/conn.go b/internal/util/socks/conn.go index 69316e8..205b585 100644 --- a/internal/util/socks/conn.go +++ b/internal/util/socks/conn.go @@ -110,7 +110,7 @@ func (c *udpConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { rbuf := bufpool.Get(c.bufferSize) defer bufpool.Put(rbuf) - n, c.raddr, err = c.PacketConn.ReadFrom(*rbuf) + n, c.raddr, err = c.PacketConn.ReadFrom(rbuf) if err != nil { return } @@ -119,11 +119,11 @@ func (c *udpConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { header := gosocks5.UDPHeader{ Addr: &socksAddr, } - hlen, err := header.ReadFrom(bytes.NewReader((*rbuf)[:n])) + hlen, err := header.ReadFrom(bytes.NewReader(rbuf[:n])) if err != nil { return } - n = copy(b, (*rbuf)[hlen:n]) + n = copy(b, rbuf[hlen:n]) addr, err = net.ResolveUDPAddr("udp", socksAddr.String()) return @@ -151,7 +151,7 @@ func (c *udpConn) WriteTo(b []byte, addr net.Addr) (n int, err error) { Data: b, } - buf := bytes.NewBuffer((*wbuf)[:0]) + buf := bytes.NewBuffer(wbuf[:0]) _, err = dgram.WriteTo(buf) if err != nil { return diff --git a/internal/util/ss/conn.go b/internal/util/ss/conn.go index 84035f3..c07cd10 100644 --- a/internal/util/ss/conn.go +++ b/internal/util/ss/conn.go @@ -45,18 +45,18 @@ func (c *UDPConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) { rbuf := bufpool.Get(c.bufferSize) defer bufpool.Put(rbuf) - n, _, err = c.PacketConn.ReadFrom(*rbuf) + n, _, err = c.PacketConn.ReadFrom(rbuf) if err != nil { return } saddr := gosocks5.Addr{} - addrLen, err := saddr.ReadFrom(bytes.NewReader((*rbuf)[:n])) + addrLen, err := saddr.ReadFrom(bytes.NewReader(rbuf[:n])) if err != nil { return } - n = copy(b, (*rbuf)[addrLen:n]) + n = copy(b, rbuf[addrLen:n]) addr, err = net.ResolveUDPAddr("udp", saddr.String()) return @@ -76,13 +76,13 @@ func (c *UDPConn) WriteTo(b []byte, addr net.Addr) (n int, err error) { return } - addrLen, err := socksAddr.Encode(*wbuf) + addrLen, err := socksAddr.Encode(wbuf) if err != nil { return } - n = copy((*wbuf)[addrLen:], b) - _, err = c.PacketConn.WriteTo((*wbuf)[:addrLen+n], c.raddr) + n = copy(wbuf[addrLen:], b) + _, err = c.PacketConn.WriteTo(wbuf[:addrLen+n], c.raddr) return } diff --git a/listener/mws/metadata.go b/listener/mws/metadata.go index 74bd365..c5d9c57 100644 --- a/listener/mws/metadata.go +++ b/listener/mws/metadata.go @@ -41,13 +41,6 @@ func (l *mwsListener) parseMetadata(md mdata.Metadata) (err error) { readBufferSize = "readBufferSize" writeBufferSize = "writeBufferSize" enableCompression = "enableCompression" - - muxKeepAliveDisabled = "muxKeepAliveDisabled" - muxKeepAliveInterval = "muxKeepAliveInterval" - muxKeepAliveTimeout = "muxKeepAliveTimeout" - muxMaxFrameSize = "muxMaxFrameSize" - muxMaxReceiveBuffer = "muxMaxReceiveBuffer" - muxMaxStreamBuffer = "muxMaxStreamBuffer" ) l.md.path = mdutil.GetString(md, path) diff --git a/listener/redirect/udp/conn.go b/listener/redirect/udp/conn.go index b023529..cae9dad 100644 --- a/listener/redirect/udp/conn.go +++ b/listener/redirect/udp/conn.go @@ -23,7 +23,7 @@ func (c *redirConn) Read(b []byte) (n int, err error) { c.once.Do(func() { n = copy(b, c.buf) - bufpool.Put(&c.buf) + bufpool.Put(c.buf) }) if n == 0 { diff --git a/listener/redirect/udp/listener_linux.go b/listener/redirect/udp/listener_linux.go index b26e500..b4c6051 100644 --- a/listener/redirect/udp/listener_linux.go +++ b/listener/redirect/udp/listener_linux.go @@ -45,7 +45,7 @@ func (l *redirectListener) listenUDP(addr string) (*net.UDPConn, error) { func (l *redirectListener) accept() (conn net.Conn, err error) { b := bufpool.Get(l.md.readBufferSize) - n, raddr, dstAddr, err := readFromUDP(l.ln, *b) + n, raddr, dstAddr, err := readFromUDP(l.ln, b) if err != nil { l.logger.Error(err) return @@ -65,7 +65,7 @@ func (l *redirectListener) accept() (conn net.Conn, err error) { conn = &redirConn{ Conn: c, - buf: (*b)[:n], + buf: b[:n], ttl: l.md.ttl, } return @@ -81,12 +81,12 @@ func readFromUDP(conn *net.UDPConn, b []byte) (n int, remoteAddr *net.UDPAddr, d oob := bufpool.Get(1024) defer bufpool.Put(oob) - n, oobn, _, remoteAddr, err := conn.ReadMsgUDP(b, *oob) + n, oobn, _, remoteAddr, err := conn.ReadMsgUDP(b, oob) if err != nil { return 0, nil, nil, err } - msgs, err := unix.ParseSocketControlMessage((*oob)[:oobn]) + msgs, err := unix.ParseSocketControlMessage(oob[:oobn]) if err != nil { return 0, nil, nil, fmt.Errorf("parsing socket control message: %s", err) } diff --git a/listener/tun/tun.go b/listener/tun/tun.go index 05d5500..22977f7 100644 --- a/listener/tun/tun.go +++ b/listener/tun/tun.go @@ -24,7 +24,7 @@ func (d *tunDevice) Read(p []byte) (n int, err error) { b := bufpool.Get(rbuf) defer bufpool.Put(b) - n, err = d.dev.Read(*b, tunOffsetBytes) + n, err = d.dev.Read(b, tunOffsetBytes) if n <= tunOffsetBytes || err != nil { d.dev.Flush() if n <= tunOffsetBytes { @@ -33,7 +33,7 @@ func (d *tunDevice) Read(p []byte) (n int, err error) { return } - n = copy(p, (*b)[tunOffsetBytes:tunOffsetBytes+n]) + n = copy(p, b[tunOffsetBytes:tunOffsetBytes+n]) return } @@ -41,8 +41,8 @@ func (d *tunDevice) Write(p []byte) (n int, err error) { b := bufpool.Get(tunOffsetBytes + len(p)) defer bufpool.Put(b) - copy((*b)[tunOffsetBytes:], p) - return d.dev.Write(*b, tunOffsetBytes) + copy(b[tunOffsetBytes:], p) + return d.dev.Write(b, tunOffsetBytes) } func (d *tunDevice) Close() error {