diff --git a/parse/addr.go b/parse/addr.go index 00cf461..e48339d 100644 --- a/parse/addr.go +++ b/parse/addr.go @@ -14,7 +14,7 @@ func ParseAddr(conn net.Conn) (teeConn *TeeConn, addr string, err error) { teeConn.StartOrReset() defer teeConn.Stop() - buf := make([]byte, 1) + buf := make([]byte, 1<<10) if n, err := teeConn.Read(buf); err != nil || n != 1 { return teeConn, "", fmt.Errorf("Read conn fail: %v, readed: %d %v", err, n, buf) } diff --git a/parse/tee_conn.go b/parse/tee_conn.go index f5b5b0c..100f6da 100644 --- a/parse/tee_conn.go +++ b/parse/tee_conn.go @@ -1,6 +1,8 @@ package parse -import "net" +import ( + "net" +) type TeeConn struct { net.Conn diff --git a/proxy/client.go b/proxy/client.go index 99dc6a2..7c9efa3 100644 --- a/proxy/client.go +++ b/proxy/client.go @@ -28,7 +28,7 @@ func StartClient(netType, server, password string) { for { conn := <-connCh - glog.V(1).Infof("new conn from (%s)", conn.RemoteAddr()) + glog.V(1).Infof("new conn from (%s) to (%s)", conn.RemoteAddr(), server) rc, err := client.Dial(server) if err != nil { diff --git a/shadow/shadow.go b/shadow/shadow.go index b2a1fbc..df948f8 100644 --- a/shadow/shadow.go +++ b/shadow/shadow.go @@ -12,22 +12,26 @@ import ( ) type conn struct { + blockSize int aead cipher.AEAD encryptNonce func() []byte decryptNonce func() []byte - readBuf []byte writeBuf []byte net.Conn } func (c *conn) Read(b []byte) (n int, err error) { + bLength := len(b) + if bLength%c.blockSize != 0 { + return 0, errors.Errorf("aead: block size %d not match", c.blockSize) + } + n, err = c.Conn.Read(b) if err != nil { return n, err } - _, err = c.aead.Open(b[:0], c.decryptNonce(), b[:n], nil) - return n - c.aead.Overhead(), err + return bLength - c.aead.Overhead(), err } func (c *conn) Write(b []byte) (n int, err error) { c.writeBuf = c.aead.Seal(nil, c.encryptNonce(), b, nil) @@ -41,20 +45,6 @@ func (c *conn) Write(b []byte) (n int, err error) { } func Shadow(c net.Conn, password string) (net.Conn, error) { - aead, err := newAEAD(password) - if err != nil { - return nil, err - } - - return &conn{ - aead: aead, - encryptNonce: newNonce(password, aead.NonceSize()), - decryptNonce: newNonce(password, aead.NonceSize()), - Conn: c, - }, nil -} - -func newAEAD(password string) (cipher.AEAD, error) { block, err := aes.NewCipher([]byte(password + password)[:16]) if err != nil { return nil, errors.Wrap(err, "password too short") @@ -65,7 +55,13 @@ func newAEAD(password string) (cipher.AEAD, error) { return nil, errors.Wrap(err, "GCM") } - return aead, nil + return &conn{ + blockSize: block.BlockSize() * 8, + aead: aead, + encryptNonce: newNonce(password, aead.NonceSize()), + decryptNonce: newNonce(password, aead.NonceSize()), + Conn: c, + }, nil } func newNonce(password string, size int) func() []byte {