Refactor crypto

This commit is contained in:
wweir
2019-01-04 16:14:15 +08:00
parent c01f217939
commit ec4286e6d4
6 changed files with 93 additions and 135 deletions
-70
View File
@@ -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
}
}
-30
View File
@@ -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)))
}
})
}
}
+5 -8
View File
@@ -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)
}
}
+9 -10
View File
@@ -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)
}
+5 -17
View File
@@ -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)
+74
View File
@@ -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
}
}