diff --git a/client/main.go b/client/main.go index 5d5ec15..9bae33a 100644 --- a/client/main.go +++ b/client/main.go @@ -20,16 +20,20 @@ import ( smuxv2 "github.com/xtaci/smux/v2" ) -// SALT is use for pbkdf2 key expansion -const SALT = "kcp-go" - -// maximum supported smux version -const maxSmuxVer = 2 +const ( + // SALT is use for pbkdf2 key expansion + SALT = "kcp-go" + // maximum supported smux version + maxSmuxVer = 2 + // stream copy buffer size + bufSize = 32768 +) // VERSION is injected by buildflags var VERSION = "SELFBUILD" -func handleClient(mux generic.Mux, p1 net.Conn, quiet bool) { +// handleClient aggregates connection p1 on mux with 'writeLock' +func handleClient(mux generic.Mux, p1 net.Conn, ctrl *generic.CopyControl, quiet bool) { logln := func(v ...interface{}) { if !quiet { log.Println(v...) @@ -53,7 +57,7 @@ func handleClient(mux generic.Mux, p1 net.Conn, quiet bool) { streamCopy := func(dst io.Writer, src io.ReadCloser) chan struct{} { die := make(chan struct{}) go func() { - if _, err := generic.Copy(dst, src); err != nil { + if _, err := generic.Copy(dst, src, ctrl); err != nil { if s2, ok := p2.(generic.Stream); ok { // verbose error handling cause := err @@ -455,11 +459,13 @@ func main() { muxes := make([]struct { session generic.Mux ttl time.Time + ctrl *generic.CopyControl // for control of memory in copying }, numconn) for k := range muxes { muxes[k].session = waitConn() muxes[k].ttl = time.Now().Add(time.Duration(config.AutoExpire) * time.Second) + muxes[k].ctrl = &generic.CopyControl{Buffer: make([]byte, bufSize)} } chScavenger := make(chan generic.Mux, 128) @@ -478,9 +484,10 @@ func main() { chScavenger <- muxes[idx].session muxes[idx].session = waitConn() muxes[idx].ttl = time.Now().Add(time.Duration(config.AutoExpire) * time.Second) + muxes[idx].ctrl = &generic.CopyControl{Buffer: make([]byte, bufSize)} } - go handleClient(muxes[idx].session, p1, config.Quiet) + go handleClient(muxes[idx].session, p1, muxes[idx].ctrl, config.Quiet) rr++ } } diff --git a/generic/copy.go b/generic/copy.go index 8a1105a..2964dcb 100644 --- a/generic/copy.go +++ b/generic/copy.go @@ -1,9 +1,20 @@ package generic -import "io" +import ( + "io" + "net" + "sync" +) -// Memory optimized Copy function specified for this library -func Copy(dst io.Writer, src io.Reader) (written int64, err error) { +const bufSize = 4096 + +type CopyControl struct { + Buffer []byte // shared buffer for copying controlled by mutex + sync.Mutex +} + +// Memory optimized io.Copy function specified for this library +func Copy(dst io.Writer, src io.Reader, ctrl *CopyControl) (written int64, err error) { // If the reader has a WriteTo method, use it to do the copy. // Avoids an allocation and a copy. if wt, ok := src.(io.WriterTo); ok { @@ -14,39 +25,16 @@ func Copy(dst io.Writer, src io.Reader) (written int64, err error) { return rt.ReadFrom(src) } - // limited to 4K per stream - size := 4096 - if l, ok := src.(*io.LimitedReader); ok && int64(size) > l.N { - if l.N < 1 { - size = 1 - } else { - size = int(l.N) + // if src is net.TCPConn, and dst is a multiplexed connection + // reading can be controlled by writable events of smux + // and make the reading serialized + if tcpconn, ok := src.(*net.TCPConn); ok { + if ctrl != nil { + return rawCopy(dst, tcpconn, ctrl) } } - buf := make([]byte, size) - for { - nr, er := src.Read(buf) - if nr > 0 { - nw, ew := dst.Write(buf[0:nr]) - if nw > 0 { - written += int64(nw) - } - if ew != nil { - err = ew - break - } - if nr != nw { - err = io.ErrShortWrite - break - } - } - if er != nil { - if er != io.EOF { - err = er - } - break - } - } - return written, err + // fallback to standard io.CopyBuffer + buf := make([]byte, bufSize) + return io.CopyBuffer(dst, src, buf) } diff --git a/generic/rawcopy_unix.go b/generic/rawcopy_unix.go new file mode 100644 index 0000000..c7f7ff6 --- /dev/null +++ b/generic/rawcopy_unix.go @@ -0,0 +1,68 @@ +// +build aix darwin dragonfly freebsd linux netbsd openbsd solaris + +package generic + +import ( + "io" + "net" + + "syscall" +) + +func rawCopy(dst io.Writer, src *net.TCPConn, ctrl *CopyControl) (written int64, err error) { + c, err := src.SyscallConn() + if err != nil { + return 0, err + } + + buf := ctrl.Buffer + for { + var er error + var nr int + rr := c.Read(func(s uintptr) bool { + ctrl.Lock() // writelock will block reading + defer ctrl.Unlock() + nr, er = syscall.Read(int(s), buf) + if er == syscall.EAGAIN { + return false + } + return true + }) + + // read EOF + if nr == 0 && er == nil { + break + } + + if nr > 0 { + ctrl.Lock() + nw, ew := dst.Write(buf[0:nr]) + ctrl.Unlock() + if nw > 0 { + written += int64(nw) + } + if ew != nil { + err = ew + break + } + if nr != nw { + err = io.ErrShortWrite + break + } + } + if er != nil { + if er != io.EOF { + err = er + } + break + } + if rr != nil { + if rr != io.EOF { + err = rr + } + break + } + } + + return written, err +} diff --git a/generic/rawcopy_windows.go b/generic/rawcopy_windows.go new file mode 100644 index 0000000..aea9686 --- /dev/null +++ b/generic/rawcopy_windows.go @@ -0,0 +1,71 @@ +// +build windows + +package generic + +import ( + "io" + "net" + "syscall" +) + +func rawCopy(dst io.Writer, src *net.TCPConn, ctrl *CopyControl) (written int64, err error) { + c, err := src.SyscallConn() + if err != nil { + return 0, err + } + + buf := ctrl.Buffer + for { + var er error + var nr int + rr := c.Read(func(s uintptr) bool { + ctrl.Lock() + defer ctrl.Unlock() + var read uint32 + var flags uint32 + var wsabuf syscall.WSABuf + wsabuf.Buf = &buf[0] + wsabuf.Len = uint32(len(buf)) + er = syscall.WSARecv(syscall.Handle(s), &wsabuf, 1, &read, &flags, nil, nil) + nr = int(read) + return true + }) + + // read EOF + if nr == 0 && er == nil { + break + } + + if nr > 0 { + ctrl.Lock() + nw, ew := dst.Write(buf[0:nr]) + ctrl.Unlock() + buf = nil + if nw > 0 { + written += int64(nw) + } + if ew != nil { + err = ew + break + } + if nr != nw { + err = io.ErrShortWrite + break + } + } + if er != nil { + if er != io.EOF { + err = er + } + break + } + if rr != nil { + if rr != io.EOF { + err = rr + } + break + } + } + + return written, err +} diff --git a/go.sum b/go.sum index f483ce9..36bea71 100644 --- a/go.sum +++ b/go.sum @@ -19,36 +19,12 @@ github.com/tjfoc/gmsm v1.0.1 h1:R11HlqhXkDospckjZEihx9SW/2VW0RgdwrykyWMFOQU= github.com/tjfoc/gmsm v1.0.1/go.mod h1:XxO4hdhhrzAd+G4CjDqaOkd0hUzmtPR/d3EiBBMn/wc= github.com/urfave/cli v1.21.0 h1:wYSSj06510qPIzGSua9ZqsncMmWE3Zr55KBERygyrxE= github.com/urfave/cli v1.21.0/go.mod h1:lxDj6qX9Q6lWQxIrbrT0nwecwUtRnhVZAJjJZrVUZZQ= -github.com/xtaci/kcp-go v5.4.13+incompatible h1:s6ba2XTw8lAj+s6AQNob25dCvWDgwE+U1QpEVBUoYy8= -github.com/xtaci/kcp-go v5.4.13+incompatible/go.mod h1:bN6vIwHQbfHaHtFpEssmWsN45a+AZwO7eyRCmEIbtvE= -github.com/xtaci/kcp-go v5.4.14+incompatible h1:kQZr/ngKQtYrgXSUxwF4A59mTMzUp0BDmtWIRuXYoqg= -github.com/xtaci/kcp-go v5.4.14+incompatible/go.mod h1:bN6vIwHQbfHaHtFpEssmWsN45a+AZwO7eyRCmEIbtvE= -github.com/xtaci/kcp-go v5.4.15+incompatible h1:QLDulPaKjT4k4cGeviyC1mt00gwJ3r5epx8yCw6ACEc= -github.com/xtaci/kcp-go v5.4.15+incompatible/go.mod h1:bN6vIwHQbfHaHtFpEssmWsN45a+AZwO7eyRCmEIbtvE= -github.com/xtaci/kcp-go v5.4.16+incompatible h1:/L7UP4P4H/oXpMnrb2W9oOxCMVpjPi3FJeMJmgN+SUE= -github.com/xtaci/kcp-go v5.4.16+incompatible/go.mod h1:bN6vIwHQbfHaHtFpEssmWsN45a+AZwO7eyRCmEIbtvE= -github.com/xtaci/kcp-go v5.4.17+incompatible h1:RudP76JCx062JSxPxSjBl+457+fS0M7T8zEZPLpC0o8= -github.com/xtaci/kcp-go v5.4.17+incompatible/go.mod h1:bN6vIwHQbfHaHtFpEssmWsN45a+AZwO7eyRCmEIbtvE= -github.com/xtaci/kcp-go v5.4.18+incompatible h1:zxzRP8V54vhJ8QAKEjf1b9g96R01prybCRchx6rEmtg= -github.com/xtaci/kcp-go v5.4.18+incompatible/go.mod h1:bN6vIwHQbfHaHtFpEssmWsN45a+AZwO7eyRCmEIbtvE= github.com/xtaci/kcp-go v5.4.19+incompatible h1:vv7Ar1D9WZGiv6deIOluxrC26Oin/2jFtx8sFU5tlvw= github.com/xtaci/kcp-go v5.4.19+incompatible/go.mod h1:bN6vIwHQbfHaHtFpEssmWsN45a+AZwO7eyRCmEIbtvE= github.com/xtaci/lossyconn v0.0.0-20190602105132-8df528c0c9ae h1:J0GxkO96kL4WF+AIT3M4mfUVinOCPgf2uUWYFUzN0sM= github.com/xtaci/lossyconn v0.0.0-20190602105132-8df528c0c9ae/go.mod h1:gXtu8J62kEgmN++bm9BVICuT/e8yiLI2KFobd/TRFsE= -github.com/xtaci/smux v1.4.4 h1:FukIfahko+KHhS9Gxppkp6756opZymvPOLNmpny1is4= -github.com/xtaci/smux v1.4.4/go.mod h1:LuA3S0xssf4fmGRJ7ow3EehgmDUzib4EcobaFNKvlMA= -github.com/xtaci/smux v1.4.5 h1:Pi64WSFFIDrpMSzRZzqt5e4YO/m4a6BAKGjxSa+KU0A= -github.com/xtaci/smux v1.4.5/go.mod h1:LuA3S0xssf4fmGRJ7ow3EehgmDUzib4EcobaFNKvlMA= github.com/xtaci/smux v1.4.6 h1:p9e/qj3Bj0zUT8qJWdmAZfmx5lOcZh0vLL0bQ8jnA7M= github.com/xtaci/smux v1.4.6/go.mod h1:LuA3S0xssf4fmGRJ7ow3EehgmDUzib4EcobaFNKvlMA= -github.com/xtaci/smux/v2 v2.0.11 h1:thVWmgGRciZ8iaATwpY2B/51aHzmMI6wrF7DfcJSckU= -github.com/xtaci/smux/v2 v2.0.11/go.mod h1:Iqy5a3Gax2p7WCKHOHkSNo/COthNFXd3/vqrcKNtzqI= -github.com/xtaci/smux/v2 v2.0.13 h1:T3VtVS8CQ2nbsMttkzTB48acij5HAhi8x8gYe9SodJA= -github.com/xtaci/smux/v2 v2.0.13/go.mod h1:Iqy5a3Gax2p7WCKHOHkSNo/COthNFXd3/vqrcKNtzqI= -github.com/xtaci/smux/v2 v2.0.14 h1:XfmTsDX7lNTLPP8bCLRQuviP1AEix3NYefM/QVUg5X4= -github.com/xtaci/smux/v2 v2.0.14/go.mod h1:Iqy5a3Gax2p7WCKHOHkSNo/COthNFXd3/vqrcKNtzqI= -github.com/xtaci/smux/v2 v2.0.15 h1:dD4yV48m6ntMVV0mieHXkxwFdx2VYnjJpqqGElXVkTs= -github.com/xtaci/smux/v2 v2.0.15/go.mod h1:Iqy5a3Gax2p7WCKHOHkSNo/COthNFXd3/vqrcKNtzqI= github.com/xtaci/smux/v2 v2.0.16 h1:2pGGbkFKTaMHIctYaovpwRpgdwWYy/6ZPaOQo00VW08= github.com/xtaci/smux/v2 v2.0.16/go.mod h1:Iqy5a3Gax2p7WCKHOHkSNo/COthNFXd3/vqrcKNtzqI= github.com/xtaci/tcpraw v1.2.25 h1:VDlqo0op17JeXBM6e2G9ocCNLOJcw9mZbobMbJjo0vk= diff --git a/server/main.go b/server/main.go index 5aff7ae..3feba5f 100644 --- a/server/main.go +++ b/server/main.go @@ -23,11 +23,14 @@ import ( "github.com/xtaci/tcpraw" ) -// SALT is use for pbkdf2 key expansion -const SALT = "kcp-go" - -// maximum supported smux version -const maxSmuxVer = 2 +const ( + // SALT is use for pbkdf2 key expansion + SALT = "kcp-go" + // maximum supported smux version + maxSmuxVer = 2 + // stream copy buffer size + bufSize = 32768 +) // VERSION is injected by buildflags var VERSION = "SELFBUILD" @@ -73,6 +76,9 @@ func handleMux(conn net.Conn, config *Config) { panic("incorrect smux version") } + // copy to stream control + copyControl := &generic.CopyControl{Buffer: make([]byte, bufSize)} + for { stream, err := muxer.Accept() if err != nil { @@ -94,12 +100,12 @@ func handleMux(conn net.Conn, config *Config) { p1.Close() return } - handleClient(p1, p2, config.Quiet) + handleClient(p1, p2, copyControl, config.Quiet) }(stream) } } -func handleClient(p1 io.ReadWriteCloser, p2 net.Conn, quiet bool) { +func handleClient(p1 io.ReadWriteCloser, p2 net.Conn, ctrl *generic.CopyControl, quiet bool) { logln := func(v ...interface{}) { if !quiet { log.Println(v...) @@ -118,7 +124,7 @@ func handleClient(p1 io.ReadWriteCloser, p2 net.Conn, quiet bool) { streamCopy := func(dst io.Writer, src io.ReadCloser) chan struct{} { die := make(chan struct{}) go func() { - if _, err := generic.Copy(dst, src); err != nil { + if _, err := generic.Copy(dst, src, ctrl); err != nil { if s1, ok := p1.(generic.Stream); ok { // verbose error handling cause := err