Add pool support everywhere

This commit is contained in:
9seconds
2020-03-24 11:22:59 +03:00
parent 9125a29e79
commit 837d96dc43
13 changed files with 162 additions and 62 deletions
+6 -1
View File
@@ -63,7 +63,12 @@ func (c *ClientProtocol) tlsHandshake(conn io.ReadWriter) error {
return fmt.Errorf("cannot read initial record: %w", err) return fmt.Errorf("cannot read initial record: %w", err)
} }
clientHello, err := tlstypes.ParseClientHello(helloRecord.Data.Bytes()) buf := acquireBytesBuffer()
defer releaseBytesBuffer(buf)
helloRecord.Data.WriteBytes(buf)
clientHello, err := tlstypes.ParseClientHello(buf.Bytes())
if err != nil { if err != nil {
return fmt.Errorf("cannot parse client hello: %w", err) return fmt.Errorf("cannot parse client hello: %w", err)
} }
+11 -8
View File
@@ -28,15 +28,9 @@ func cloak(one, another io.ReadWriteCloser) {
wg.Add(2) wg.Add(2)
go func() { go cloakPipe(one, another, wg)
defer wg.Done()
io.Copy(one, another) // nolint: errcheck
}()
go func() { go cloakPipe(another, one, wg)
defer wg.Done()
io.Copy(another, one) // nolint: errcheck
}()
go func() { go func() {
wg.Wait() wg.Wait()
@@ -69,3 +63,12 @@ func cloak(one, another io.ReadWriteCloser) {
<-ctx.Done() <-ctx.Done()
} }
func cloakPipe(one io.Writer, another io.Reader, wg *sync.WaitGroup) {
defer wg.Done()
buf := acquireCloakBuffer()
defer releaseCloakBuffer(buf)
io.CopyBuffer(one, another, *buf) // nolint: errcheck
}
+38
View File
@@ -0,0 +1,38 @@
package faketls
import (
"bytes"
"sync"
)
const cloakBufferSize = 1024
var (
poolBytesBuffer = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
poolCloakBuffer = sync.Pool{
New: func() interface{} {
rv := make([]byte, cloakBufferSize)
return &rv
},
}
)
func acquireBytesBuffer() *bytes.Buffer {
return poolBytesBuffer.Get().(*bytes.Buffer)
}
func acquireCloakBuffer() *[]byte {
return poolCloakBuffer.Get().(*[]byte)
}
func releaseBytesBuffer(buf *bytes.Buffer) {
poolBytesBuffer.Put(buf)
}
func releaseCloakBuffer(buf *[]byte) {
poolCloakBuffer.Put(buf)
}
+1 -1
View File
@@ -25,7 +25,7 @@ func (c ClientHello) Digest() []byte {
} }
mac := hmac.New(sha256.New, config.C.Secret) mac := hmac.New(sha256.New, config.C.Secret)
mac.Write(rec.Bytes()) // nolint: errcheck rec.WriteBytes(mac)
computedDigest := mac.Sum(nil) computedDigest := mac.Sum(nil)
for i := range computedDigest { for i := range computedDigest {
+10 -3
View File
@@ -1,5 +1,7 @@
package tlstypes package tlstypes
import "io"
type RecordType uint8 type RecordType uint8
const ( const (
@@ -69,11 +71,16 @@ var (
) )
type Byter interface { type Byter interface {
Bytes() []byte WriteBytes(io.Writer)
Len() int
} }
type RawBytes []byte type RawBytes []byte
func (r RawBytes) Bytes() []byte { func (r RawBytes) WriteBytes(writer io.Writer) {
return []byte(r) writer.Write(r) // nolint: errcheck
}
func (r RawBytes) Len() int {
return len(r)
} }
+16 -9
View File
@@ -1,7 +1,7 @@
package tlstypes package tlstypes
import ( import (
"bytes" "io"
"github.com/9seconds/mtg/utils" "github.com/9seconds/mtg/utils"
) )
@@ -14,24 +14,31 @@ type Handshake struct {
Tail Byter Tail Byter
} }
func (h *Handshake) Bytes() []byte { func (h *Handshake) WriteBytes(writer io.Writer) {
buf := bytes.Buffer{} packetBuf := acquireBytesBuffer()
packetBuf := bytes.Buffer{} defer releaseBytesBuffer(packetBuf)
buf.WriteByte(byte(h.Type)) writer.Write([]byte{byte(h.Type)}) // nolint: errcheck
packetBuf.Write(h.Version.Bytes()) packetBuf.Write(h.Version.Bytes())
packetBuf.Write(h.Random[:]) packetBuf.Write(h.Random[:])
packetBuf.WriteByte(byte(len(h.SessionID))) packetBuf.WriteByte(byte(len(h.SessionID)))
packetBuf.Write(h.SessionID) packetBuf.Write(h.SessionID)
packetBuf.Write(h.Tail.Bytes()) h.Tail.WriteBytes(packetBuf)
sizeUint24 := utils.ToUint24(uint32(packetBuf.Len())) sizeUint24 := utils.ToUint24(uint32(packetBuf.Len()))
sizeUint24Bytes := sizeUint24[:] sizeUint24Bytes := sizeUint24[:]
sizeUint24Bytes[0], sizeUint24Bytes[2] = sizeUint24Bytes[2], sizeUint24Bytes[0] sizeUint24Bytes[0], sizeUint24Bytes[2] = sizeUint24Bytes[2], sizeUint24Bytes[0]
buf.Write(sizeUint24Bytes) writer.Write(sizeUint24Bytes) // nolint: errcheck
packetBuf.WriteTo(&buf) // nolint: errcheck packetBuf.WriteTo(writer) // nolint: errcheck
}
return buf.Bytes() func (h *Handshake) Len() int {
buf := acquireBytesBuffer()
defer releaseBytesBuffer(buf)
h.WriteBytes(buf)
return buf.Len()
} }
+23
View File
@@ -0,0 +1,23 @@
package tlstypes
import (
"bytes"
"sync"
)
var (
poolBytesBuffer = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func acquireBytesBuffer() *bytes.Buffer {
return poolBytesBuffer.Get().(*bytes.Buffer)
}
func releaseBytesBuffer(buf *bytes.Buffer) {
buf.Reset()
poolBytesBuffer.Put(buf)
}
+8 -9
View File
@@ -15,16 +15,15 @@ type Record struct {
Data Byter Data Byter
} }
func (r Record) Bytes() []byte { func (r Record) WriteBytes(writer io.Writer) {
buf := bytes.Buffer{} writer.Write([]byte{byte(r.Type)}) // nolint: errcheck
data := r.Data.Bytes() writer.Write(r.Version.Bytes()) // nolint: errcheck
binary.Write(writer, binary.BigEndian, uint16(r.Data.Len())) // nolint: errcheck
r.Data.WriteBytes(writer)
}
buf.WriteByte(byte(r.Type)) func (r Record) Len() int {
buf.Write(r.Version.Bytes()) return 1 + 2 + 2 + r.Data.Len()
binary.Write(&buf, binary.BigEndian, uint16(len(data))) // nolint: errcheck
buf.Write(data)
return buf.Bytes()
} }
func ReadRecord(reader io.Reader) (Record, error) { func ReadRecord(reader io.Reader) (Record, error) {
+6 -3
View File
@@ -20,20 +20,22 @@ type ServerHello struct {
} }
func (s ServerHello) WelcomePacket() []byte { func (s ServerHello) WelcomePacket() []byte {
buf := &bytes.Buffer{}
s.Random = [32]byte{} s.Random = [32]byte{}
rec := Record{ rec := Record{
Type: RecordTypeHandshake, Type: RecordTypeHandshake,
Version: Version12, Version: Version12,
Data: &s, Data: &s,
} }
buf := bytes.NewBuffer(rec.Bytes()) rec.WriteBytes(buf)
recChangeCipher := Record{ recChangeCipher := Record{
Type: RecordTypeChangeCipherSpec, Type: RecordTypeChangeCipherSpec,
Version: Version12, Version: Version12,
Data: RawBytes([]byte{0x01}), Data: RawBytes([]byte{0x01}),
} }
buf.Write(recChangeCipher.Bytes()) recChangeCipher.WriteBytes(buf)
hostCert := make([]byte, 1024+mrand.Intn(3092)) hostCert := make([]byte, 1024+mrand.Intn(3092))
rand.Read(hostCert) // nolint: errcheck rand.Read(hostCert) // nolint: errcheck
@@ -43,7 +45,8 @@ func (s ServerHello) WelcomePacket() []byte {
Version: Version12, Version: Version12,
Data: RawBytes(hostCert), Data: RawBytes(hostCert),
} }
buf.Write(recData.Bytes()) recData.WriteBytes(buf)
packet := buf.Bytes() packet := buf.Bytes()
mac := hmac.New(sha256.New, config.C.Secret) mac := hmac.New(sha256.New, config.C.Secret)
+13 -3
View File
@@ -1,6 +1,7 @@
package stream package stream
import ( import (
"bytes"
"errors" "errors"
"fmt" "fmt"
"net" "net"
@@ -39,13 +40,19 @@ func (w *wrapperFakeTLS) WriteTimeout(p []byte, timeout time.Duration) (int, err
func (w *wrapperFakeTLS) write(p []byte, writeFunc func([]byte) (int, error)) (int, error) { func (w *wrapperFakeTLS) write(p []byte, writeFunc func([]byte) (int, error)) (int, error) {
sum := 0 sum := 0
buf := acquireBytesBuffer()
defer releaseBytesBuffer(buf)
for _, v := range tlstypes.MakeRecords(p) { for _, v := range tlstypes.MakeRecords(p) {
_, err := writeFunc(v.Bytes()) buf.Reset()
v.WriteBytes(buf)
_, err := writeFunc(buf.Bytes())
if err != nil { if err != nil {
return sum, err return sum, err
} }
sum += len(v.Data.Bytes()) sum += v.Data.Len()
} }
return sum, nil return sum, nil
@@ -86,7 +93,10 @@ func NewFakeTLS(socket conntypes.StreamReadWriteCloser) conntypes.StreamReadWrit
switch rec.Type { switch rec.Type {
case tlstypes.RecordTypeChangeCipherSpec: case tlstypes.RecordTypeChangeCipherSpec:
case tlstypes.RecordTypeApplicationData: case tlstypes.RecordTypeApplicationData:
return rec.Data.Bytes(), nil buf := &bytes.Buffer{}
rec.Data.WriteBytes(buf)
return buf.Bytes(), nil
default: default:
return nil, fmt.Errorf("unsupported record type %v", rec.Type) return nil, fmt.Errorf("unsupported record type %v", rec.Type)
} }
+3 -2
View File
@@ -1,7 +1,6 @@
package stream package stream
import ( import (
"bytes"
"crypto/aes" "crypto/aes"
"crypto/cipher" "crypto/cipher"
"crypto/md5" // nolint: gosec "crypto/md5" // nolint: gosec
@@ -54,7 +53,9 @@ func mtprotoDeriveKeys(purpose mtprotoCipherPurpose,
resp *rpc.NonceResponse, resp *rpc.NonceResponse,
client, remote *net.TCPAddr, client, remote *net.TCPAddr,
secret []byte) ([]byte, []byte) { secret []byte) ([]byte, []byte) {
message := bytes.Buffer{} message := acquireBytesBuffer()
defer releaseBytesBuffer(message)
message.Write(resp.Nonce) // nolint: gosec message.Write(resp.Nonce) // nolint: gosec
message.Write(req.Nonce) // nolint: gosec message.Write(req.Nonce) // nolint: gosec
message.Write(req.CryptoTS) // nolint: gosec message.Write(req.CryptoTS) // nolint: gosec
+4 -23
View File
@@ -1,11 +1,9 @@
package stream package stream
import ( import (
"bytes"
"crypto/cipher" "crypto/cipher"
"fmt" "fmt"
"net" "net"
"sync"
"time" "time"
"go.uber.org/zap" "go.uber.org/zap"
@@ -13,23 +11,6 @@ import (
"github.com/9seconds/mtg/conntypes" "github.com/9seconds/mtg/conntypes"
) )
var (
poolWrapperObfuscated2WritePool = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func poolWrapperObfuscated2WritePoolAcquire() *bytes.Buffer {
return poolWrapperObfuscated2WritePool.Get().(*bytes.Buffer)
}
func poolWrapperObfuscated2WritePoolRelease(buf *bytes.Buffer) {
buf.Reset()
poolWrapperObfuscated2WritePool.Put(buf)
}
type wrapperObfuscated2 struct { type wrapperObfuscated2 struct {
encryptor cipher.Stream encryptor cipher.Stream
decryptor cipher.Stream decryptor cipher.Stream
@@ -59,8 +40,8 @@ func (w *wrapperObfuscated2) Read(p []byte) (int, error) {
} }
func (w *wrapperObfuscated2) WriteTimeout(p []byte, timeout time.Duration) (int, error) { func (w *wrapperObfuscated2) WriteTimeout(p []byte, timeout time.Duration) (int, error) {
buffer := poolWrapperObfuscated2WritePoolAcquire() buffer := acquireBytesBuffer()
defer poolWrapperObfuscated2WritePoolRelease(buffer) defer releaseBytesBuffer(buffer)
buffer.Write(p) buffer.Write(p)
@@ -72,8 +53,8 @@ func (w *wrapperObfuscated2) WriteTimeout(p []byte, timeout time.Duration) (int,
} }
func (w *wrapperObfuscated2) Write(p []byte) (int, error) { func (w *wrapperObfuscated2) Write(p []byte) (int, error) {
buffer := poolWrapperObfuscated2WritePoolAcquire() buffer := acquireBytesBuffer()
defer poolWrapperObfuscated2WritePoolRelease(buffer) defer releaseBytesBuffer(buffer)
buffer.Write(p) buffer.Write(p)
+23
View File
@@ -0,0 +1,23 @@
package stream
import (
"bytes"
"sync"
)
var (
poolBytesBuffer = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func acquireBytesBuffer() *bytes.Buffer {
return poolBytesBuffer.Get().(*bytes.Buffer)
}
func releaseBytesBuffer(buf *bytes.Buffer) {
buf.Reset()
poolBytesBuffer.Put(buf)
}