diff --git a/crypto/crypto.go b/crypto/crypto.go deleted file mode 100644 index 841345c..0000000 --- a/crypto/crypto.go +++ /dev/null @@ -1,70 +0,0 @@ -package crypto - -import ( - "crypto/aes" - "crypto/cipher" - "encoding/binary" - "math/rand" - - "github.com/golang/glog" - "github.com/pkg/errors" -) - -type Crypto struct { - password string - cipher.AEAD -} - -func NewCrypto(password string) (*Crypto, error) { - aead, err := newAEAD(password) - if err != nil { - return nil, errors.Wrap(err, "AEAD") - } - - return &Crypto{ - password: password, - AEAD: aead, - }, nil -} - -func (c *Crypto) Crypto() (encrypt, decrypt func(src []byte) []byte) { - nonceEncrypt := newNonce(c.password, c.AEAD.NonceSize()) - nonceDecrypt := newNonce(c.password, c.AEAD.NonceSize()) - - return func(src []byte) []byte { - return c.AEAD.Seal(nil, nonceEncrypt(), src, nil) - }, - func(src []byte) []byte { - dst, err := c.AEAD.Open(nil, nonceDecrypt(), src, nil) - if err != nil { - glog.Fatalf("%+v", errors.Wrap(err, "decrypt")) - } - - return dst - } -} - -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") - } - - aead, err := cipher.NewGCM(block) - if err != nil { - return nil, errors.Wrap(err, "GCM") - } - - return aead, nil -} - -func newNonce(password string, size int) func() []byte { - num, _ := binary.Varint([]byte(password)) - rnd := rand.New(rand.NewSource(num)) - - buf := make([]byte, size) - return func() []byte { - rnd.Read(buf) - return buf - } -} diff --git a/crypto/crypto_test.go b/crypto/crypto_test.go deleted file mode 100644 index f5556af..0000000 --- a/crypto/crypto_test.go +++ /dev/null @@ -1,30 +0,0 @@ -package crypto - -import ( - "reflect" - "testing" -) - -func TestCrypto_Crypto(t *testing.T) { - mockCrypto, err := NewCrypto("12345678") - if err != nil { - t.Errorf("%s", err) - } - - tests := []struct { - name string - data []byte - }{{ - "", - []byte("123"), - }} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - gotEncrypt, gotDecrypt := mockCrypto.Crypto() - - if !reflect.DeepEqual(tt.data, gotDecrypt(gotEncrypt(tt.data))) { - t.Errorf("Crypto.Crypto() raw = %v, got = %v", tt.data, gotDecrypt(gotEncrypt(tt.data))) - } - }) - } -} diff --git a/proxy/client.go b/proxy/client.go index 273d3b7..99dc6a2 100644 --- a/proxy/client.go +++ b/proxy/client.go @@ -4,10 +4,10 @@ import ( "net" "github.com/golang/glog" - "github.com/wweir/sower/crypto" "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 { @@ -26,11 +26,6 @@ func StartClient(netType, server, password string) { client = tcp.NewClient() } - cryptor, err := crypto.NewCrypto(password) - if err != nil { - glog.Fatalln(err) - } - for { conn := <-connCh glog.V(1).Infof("new conn from (%s)", conn.RemoteAddr()) @@ -42,9 +37,11 @@ func StartClient(netType, server, password string) { continue } - encrypt, decrypt := cryptor.Crypto() + if rc, err = shadow.Shadow(rc, password); err != nil { + glog.Fatalln(err) + } - go relay(conn, rc, encrypt, decrypt) + go relay(conn, rc) } } diff --git a/proxy/server.go b/proxy/server.go index b7a1cff..e107538 100644 --- a/proxy/server.go +++ b/proxy/server.go @@ -5,11 +5,11 @@ import ( "strings" "github.com/golang/glog" - "github.com/wweir/sower/crypto" "github.com/wweir/sower/parse" "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 { @@ -38,18 +38,18 @@ func StartServer(netType, port, password string) { glog.Fatalf("listen %v fail: %s", port, err) } - cryptor, err := crypto.NewCrypto(password) - if err != nil { - glog.Fatalln(err) - } - for { conn := <-connCh - go handle(conn, cryptor) + conn, err := shadow.Shadow(conn, password) + if err != nil { + glog.Fatalln(err) + } + + go handle(conn) } } -func handle(conn net.Conn, cryptor *crypto.Crypto) { +func handle(conn net.Conn) { defer conn.Close() conn, addr, err := parse.ParseAddr(conn) @@ -69,6 +69,5 @@ func handle(conn net.Conn, cryptor *crypto.Crypto) { glog.Warningln(err) } - encrypt, decrypt := cryptor.Crypto() - relay(rc, conn, encrypt, decrypt) + relay(rc, conn) } diff --git a/proxy/util.go b/proxy/util.go index 63cd3d9..962ea9a 100644 --- a/proxy/util.go +++ b/proxy/util.go @@ -19,29 +19,17 @@ const ( TCP ) -func relay(transparentConn, cryptoConn net.Conn, encrypt, decrypt func([]byte) []byte) { +func relay(conn1, conn2 net.Conn) { wg := &sync.WaitGroup{} exitFlag := new(int32) wg.Add(2) - go redirect(cryptoConn, transparentConn, encrypt, wg, exitFlag) - redirect(transparentConn, cryptoConn, decrypt, wg, exitFlag) + go redirect(conn2, conn1, wg, exitFlag) + redirect(conn1, conn2, wg, exitFlag) wg.Wait() } -func redirect(dst, src net.Conn, fn func([]byte) []byte, wg *sync.WaitGroup, exitFlag *int32) { - var buf = make([]byte, 4<<20 /*4M*/) - var n int - var err error - for { - if n, err = src.Read(buf); err != nil { - break - } - if _, err = dst.Write(fn(buf[:n])); err != nil { - break - } - } - - if err != io.EOF && (atomic.LoadInt32(exitFlag) == 0) { +func redirect(dst, src net.Conn, wg *sync.WaitGroup, exitFlag *int32) { + if _, err := io.Copy(dst, src); err != io.EOF && (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 new file mode 100644 index 0000000..7199e7f --- /dev/null +++ b/shadow/shadow.go @@ -0,0 +1,74 @@ +package shadow + +import ( + "crypto/aes" + "crypto/cipher" + "encoding/binary" + "math/rand" + "net" + + "github.com/pkg/errors" +) + +type conn struct { + 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) { + n, err = c.Conn.Read(b) + if err != nil { + return n, err + } + + c.readBuf, err = c.aead.Open(b[:0], c.decryptNonce(), b[:n], nil) + return len(c.readBuf), err +} + +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 (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) +} + +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") + } + + aead, err := cipher.NewGCM(block) + if err != nil { + return nil, errors.Wrap(err, "GCM") + } + + return aead, nil +} + +func newNonce(password string, size int) func() []byte { + num, _ := binary.Varint([]byte(password)) + rnd := rand.New(rand.NewSource(num)) + + buf := make([]byte, size) + return func() []byte { + rnd.Read(buf) + return buf + } +}