diff --git a/client/main.go b/client/main.go index fb55b37..c026442 100644 --- a/client/main.go +++ b/client/main.go @@ -18,6 +18,7 @@ import ( "github.com/pkg/errors" "github.com/urfave/cli" kcp "github.com/xtaci/kcp-go" + "github.com/xtaci/kcptun/generic" "github.com/xtaci/smux" "path/filepath" @@ -86,21 +87,11 @@ func handleClient(sess *smux.Session, p1 io.ReadWriteCloser, quiet bool) { wg.Add(2) // start tunnel & wait for tunnel termination streamCopy := func(dst io.Writer, src io.ReadCloser) { - if wt, ok := src.(io.WriterTo); ok { - if _, err := wt.WriteTo(dst); err != nil { - logln(err) - } - } else if rt, ok := dst.(io.ReaderFrom); ok { - if _, err := rt.ReadFrom(src); err != nil { - logln(err) - } - } else { - buf := xmitBuf.Get().([]byte) - if _, err := io.CopyBuffer(dst, src, buf); err != nil { - logln(err) - } - xmitBuf.Put(buf) + buf := xmitBuf.Get().([]byte) + if _, err := generic.CopyBuffer(dst, src, buf); err != nil { + logln(err) } + xmitBuf.Put(buf) src.Close() wg.Done() } diff --git a/generic/copy.go b/generic/copy.go new file mode 100644 index 0000000..023556b --- /dev/null +++ b/generic/copy.go @@ -0,0 +1,36 @@ +package generic + +import "io" + +// io.CopyBuffer has extra tests for interface like io.ReaderFrom and io.WriterTo +// which is not efficient in memory management from tests +func CopyBuffer(dst io.Writer, src io.Reader, buf []byte) (written int64, err error) { + if buf != nil && len(buf) == 0 { + panic("empty buffer in copyBuffer") + } + + 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 +} diff --git a/server/main.go b/server/main.go index e55cd26..659927b 100644 --- a/server/main.go +++ b/server/main.go @@ -22,6 +22,7 @@ import ( "github.com/pkg/errors" "github.com/urfave/cli" kcp "github.com/xtaci/kcp-go" + "github.com/xtaci/kcptun/generic" "github.com/xtaci/smux" ) @@ -113,21 +114,11 @@ func handleClient(p1 *smux.Stream, p2 io.ReadWriteCloser, quiet bool) { wg.Add(2) // start tunnel & wait for tunnel termination streamCopy := func(dst io.Writer, src io.ReadCloser) { - if wt, ok := src.(io.WriterTo); ok { - if _, err := wt.WriteTo(dst); err != nil { - logln(err) - } - } else if rt, ok := dst.(io.ReaderFrom); ok { - if _, err := rt.ReadFrom(src); err != nil { - logln(err) - } - } else { - buf := xmitBuf.Get().([]byte) - if _, err := io.CopyBuffer(dst, src, buf); err != nil { - logln(err) - } - xmitBuf.Put(buf) + buf := xmitBuf.Get().([]byte) + if _, err := generic.CopyBuffer(dst, src, buf); err != nil { + logln(err) } + xmitBuf.Put(buf) src.Close() wg.Done() }