diff --git a/proxy/client.go b/proxy/client.go index 08f701b..99dc6a2 100644 --- a/proxy/client.go +++ b/proxy/client.go @@ -7,6 +7,7 @@ import ( "github.com/wweir/sower/proxy/kcp" "github.com/wweir/sower/proxy/quic" "github.com/wweir/sower/proxy/tcp" + "github.com/wweir/sower/shadow" ) type Client interface { @@ -36,9 +37,9 @@ func StartClient(netType, server, password string) { continue } - // if rc, err = shadow.Shadow(rc, password); err != nil { - // glog.Fatalln(err) - // } + if rc, err = shadow.Shadow(rc, password); err != nil { + glog.Fatalln(err) + } go relay(conn, rc) } diff --git a/proxy/server.go b/proxy/server.go index d6646be..f367cf9 100644 --- a/proxy/server.go +++ b/proxy/server.go @@ -9,6 +9,7 @@ import ( "github.com/wweir/sower/proxy/kcp" "github.com/wweir/sower/proxy/quic" "github.com/wweir/sower/proxy/tcp" + "github.com/wweir/sower/shadow" ) type Server interface { @@ -39,10 +40,10 @@ func StartServer(netType, port, password string) { for { conn := <-connCh - // conn, err := shadow.Shadow(conn, password) - // if err != nil { - // glog.Fatalln(err) - // } + conn, err := shadow.Shadow(conn, password) + if err != nil { + glog.Fatalln(err) + } go handle(conn) } diff --git a/proxy/util.go b/proxy/util.go index f427889..c3cb827 100644 --- a/proxy/util.go +++ b/proxy/util.go @@ -28,7 +28,7 @@ func relay(conn1, conn2 net.Conn) { } func redirect(dst, src net.Conn, wg *sync.WaitGroup, exitFlag *int32) { - if _, err := io.Copy(dst, src); err != io.EOF && (atomic.LoadInt32(exitFlag) == 0) { + if _, err := io.Copy(dst, src); err != nil && (atomic.LoadInt32(exitFlag) == 0) { glog.V(1).Infof("%s<>%s -> %s<>%s: %s", src.RemoteAddr(), src.LocalAddr(), dst.LocalAddr(), dst.RemoteAddr(), err) } atomic.AddInt32(exitFlag, 1) diff --git a/shadow/shadow.go b/shadow/shadow.go index 2e4bc41..b2a1fbc 100644 --- a/shadow/shadow.go +++ b/shadow/shadow.go @@ -4,6 +4,7 @@ import ( "crypto/aes" "crypto/cipher" "encoding/binary" + "io" "math/rand" "net" @@ -25,12 +26,18 @@ func (c *conn) Read(b []byte) (n int, err error) { return n, err } - c.readBuf, err = c.aead.Open(b[:0], c.decryptNonce(), b[:n], nil) - return len(c.readBuf), err + _, err = c.aead.Open(b[:0], c.decryptNonce(), b[:n], nil) + return n - c.aead.Overhead(), err } func (c *conn) Write(b []byte) (n int, err error) { c.writeBuf = c.aead.Seal(nil, c.encryptNonce(), b, nil) - return c.Conn.Write(c.writeBuf) + for n < len(c.writeBuf) { + n, err = c.Conn.Write(c.writeBuf) + if err != nil && err != io.EOF { + return 0, err + } + } + return len(b), err } func Shadow(c net.Conn, password string) (net.Conn, error) {