From c62e3a3284cb2f1a1dc0095d225e249bdcd211e1 Mon Sep 17 00:00:00 2001 From: wweir Date: Fri, 4 Jan 2019 14:34:00 +0800 Subject: [PATCH] Add crypto support --- crypto/crypto.go | 69 +++++++++++++++++++++++++++++++++++++++++ proxy/client.go | 12 +++++-- proxy/kcp/client.go | 9 ++---- proxy/kcp/server.go | 8 ++--- proxy/kcp/util.go | 8 ----- proxy/nettype_string.go | 4 +-- proxy/quic/client.go | 7 ++--- proxy/quic/server.go | 3 +- proxy/server.go | 14 +++++++-- proxy/tcp/server.go | 3 +- proxy/util.go | 30 ++++++++++++------ 11 files changed, 123 insertions(+), 44 deletions(-) create mode 100644 crypto/crypto.go delete mode 100644 proxy/kcp/util.go diff --git a/crypto/crypto.go b/crypto/crypto.go new file mode 100644 index 0000000..5efcc41 --- /dev/null +++ b/crypto/crypto.go @@ -0,0 +1,69 @@ +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) { + nonce := newNonce(c.password, c.AEAD.NonceSize()) + + return func(src []byte) []byte { + return c.AEAD.Seal(nil, nonce(), src, nil) + }, + func(src []byte) []byte { + dst, err := c.AEAD.Open(nil, nonce(), src, nil) + if err != nil { + glog.Fatalln(err) + } + + 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/proxy/client.go b/proxy/client.go index e3f0902..273d3b7 100644 --- a/proxy/client.go +++ b/proxy/client.go @@ -4,6 +4,7 @@ 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" @@ -20,11 +21,16 @@ func StartClient(netType, server, password string) { case QUIC.String(): client = quic.NewClient() case KCP.String(): - client = kcp.NewClient(password) + client = kcp.NewClient() case TCP.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()) @@ -36,7 +42,9 @@ func StartClient(netType, server, password string) { continue } - go relay(conn, rc) + encrypt, decrypt := cryptor.Crypto() + + go relay(conn, rc, encrypt, decrypt) } } diff --git a/proxy/kcp/client.go b/proxy/kcp/client.go index d54cb23..8092e43 100644 --- a/proxy/kcp/client.go +++ b/proxy/kcp/client.go @@ -8,7 +8,6 @@ import ( ) type client struct { - Password []byte DataShard int ParityShard int DSCP int @@ -23,9 +22,8 @@ type client struct { MTU int } -func NewClient(password string) *client { +func NewClient() *client { return &client{ - Password: fillPassword(password), DataShard: 10, ParityShard: 3, DSCP: 0, @@ -41,10 +39,9 @@ func NewClient(password string) *client { } func (c *client) Dial(server string) (net.Conn, error) { - block, _ := kcp.NewAESBlockCrypt(c.Password) - conn, err := kcp.DialWithOptions(server, block, c.DataShard, c.ParityShard) + conn, err := kcp.DialWithOptions(server, nil, c.DataShard, c.ParityShard) if err != nil { - return nil, errors.Wrap(err, "createConn()") + return nil, errors.Wrap(err, "dial") } conn.SetStreamMode(true) diff --git a/proxy/kcp/server.go b/proxy/kcp/server.go index dde60b8..4156cd0 100644 --- a/proxy/kcp/server.go +++ b/proxy/kcp/server.go @@ -9,7 +9,6 @@ import ( ) type server struct { - Password []byte DataShard int ParityShard int DSCP int @@ -18,7 +17,6 @@ type server struct { func NewServer(password string) *server { return &server{ - Password: fillPassword(password), DataShard: 10, ParityShard: 3, DSCP: 0, @@ -27,8 +25,7 @@ func NewServer(password string) *server { } func (s *server) Listen(port string) (<-chan net.Conn, error) { - block, _ := kcp.NewAESBlockCrypt(s.Password) - ln, err := kcp.ListenWithOptions(port, block, s.DataShard, s.ParityShard) + ln, err := kcp.ListenWithOptions(port, nil, s.DataShard, s.ParityShard) if err != nil { return nil, err } @@ -48,8 +45,7 @@ func (s *server) Listen(port string) (<-chan net.Conn, error) { for { conn, err := ln.AcceptKCP() if err != nil { - glog.Errorln(err) - continue + glog.Fatalln("KCP listen:", err) } connCh <- conn diff --git a/proxy/kcp/util.go b/proxy/kcp/util.go deleted file mode 100644 index 4b2ef0b..0000000 --- a/proxy/kcp/util.go +++ /dev/null @@ -1,8 +0,0 @@ -package kcp - -func fillPassword(password string) []byte { - for len(password) < 16 { - password += password - } - return []byte(password)[:16] -} diff --git a/proxy/nettype_string.go b/proxy/nettype_string.go index 7641bf5..fed04b6 100644 --- a/proxy/nettype_string.go +++ b/proxy/nettype_string.go @@ -4,9 +4,9 @@ package proxy import "strconv" -const _netType_name = "QUICKCP" +const _netType_name = "QUICKCPTCP" -var _netType_index = [...]uint8{0, 4, 7} +var _netType_index = [...]uint8{0, 4, 7, 10} func (i netType) String() string { if i < 0 || i >= netType(len(_netType_index)-1) { diff --git a/proxy/quic/client.go b/proxy/quic/client.go index db740bf..be7e16b 100644 --- a/proxy/quic/client.go +++ b/proxy/quic/client.go @@ -10,9 +10,8 @@ import ( ) type client struct { - server string - conf *quic.Config - sess quic.Session + conf *quic.Config + sess quic.Session } func NewClient() *client { @@ -27,7 +26,7 @@ func NewClient() *client { func (c *client) Dial(server string) (net.Conn, error) { if c.sess == nil { - if sess, err := quic.DialAddr(c.server, &tls.Config{InsecureSkipVerify: true}, c.conf); err != nil { + if sess, err := quic.DialAddr(server, &tls.Config{InsecureSkipVerify: true}, c.conf); err != nil { return nil, errors.Wrap(err, "session") } else { c.sess = sess diff --git a/proxy/quic/server.go b/proxy/quic/server.go index 198d526..309d736 100644 --- a/proxy/quic/server.go +++ b/proxy/quic/server.go @@ -50,8 +50,7 @@ func accept(sess quic.Session, connCh chan<- net.Conn) { for { stream, err := sess.AcceptStream() if err != nil { - glog.Errorln(err) - return + glog.Fatalln("QUIC listen:", err) } connCh <- &streamConn{stream, sess} diff --git a/proxy/server.go b/proxy/server.go index 124fc10..b7a1cff 100644 --- a/proxy/server.go +++ b/proxy/server.go @@ -5,6 +5,7 @@ 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" @@ -37,13 +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) + go handle(conn, cryptor) } } -func handle(conn net.Conn) { +func handle(conn net.Conn, cryptor *crypto.Crypto) { defer conn.Close() conn, addr, err := parse.ParseAddr(conn) @@ -62,5 +68,7 @@ func handle(conn net.Conn) { if err := rc.(*net.TCPConn).SetKeepAlive(true); err != nil { glog.Warningln(err) } - relay(rc, conn) + + encrypt, decrypt := cryptor.Crypto() + relay(rc, conn, encrypt, decrypt) } diff --git a/proxy/tcp/server.go b/proxy/tcp/server.go index b5d2b65..1cc181c 100644 --- a/proxy/tcp/server.go +++ b/proxy/tcp/server.go @@ -24,8 +24,7 @@ func (s *server) Listen(port string) (<-chan net.Conn, error) { for { conn, err := ln.Accept() if err != nil { - glog.Errorln(err) - continue + glog.Fatalln("TCP listen:", err) } connCh <- conn diff --git a/proxy/util.go b/proxy/util.go index c1b16f1..5c3eda1 100644 --- a/proxy/util.go +++ b/proxy/util.go @@ -19,24 +19,36 @@ const ( TCP ) -func relay(conn1, conn2 net.Conn) { +func relay(transparentConn, cryptoConn net.Conn, encrypt, decrypt func([]byte) []byte) { wg := &sync.WaitGroup{} exitFlag := new(int32) wg.Add(2) - go redirect(conn1, conn2, wg, exitFlag) - redirect(conn2, conn1, wg, exitFlag) + go redirect(transparentConn, cryptoConn, encrypt, wg, exitFlag) + redirect(cryptoConn, transparentConn, decrypt, wg, exitFlag) wg.Wait() } -func redirect(conn1, conn2 net.Conn, wg *sync.WaitGroup, exitFlag *int32) { - if _, err := io.Copy(conn2, conn1); err != nil && (atomic.LoadInt32(exitFlag) == 0) { - glog.V(1).Infof("%s<>%s -> %s<>%s: %s", conn1.RemoteAddr(), conn1.LocalAddr(), conn2.LocalAddr(), conn2.RemoteAddr(), err) +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 + } } - // wakeup all conn goroutine + if 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) + + // wakeup all conn goroutine now := time.Now() - conn1.SetDeadline(now) - conn2.SetDeadline(now) + dst.SetDeadline(now) + src.SetDeadline(now) wg.Done() }