mirror of
https://github.com/wweir/sower.git
synced 2024-04-21 12:42:15 +00:00
Refactor crypto
This commit is contained in:
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user