Merge pull request #140 from 9seconds/pools

Add pool support everywhere
This commit is contained in:
Sergey Arkhipov
2020-03-24 16:53:48 +03:00
committed by GitHub
19 changed files with 223 additions and 71 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
}
+39
View File
@@ -0,0 +1,39 @@
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) {
buf.Reset()
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)
+5 -4
View File
@@ -42,7 +42,9 @@ type wrapperMtprotoFrame struct {
} }
func (w *wrapperMtprotoFrame) Read() (conntypes.Packet, error) { // nolint: funlen func (w *wrapperMtprotoFrame) Read() (conntypes.Packet, error) { // nolint: funlen
buf := &bytes.Buffer{} buf := acquireMtprotoFrameBytesBuffer()
defer releaseMtprotoFrameBytesBuffer(buf)
sum := crc32.NewIEEE() sum := crc32.NewIEEE()
writer := io.MultiWriter(buf, sum) writer := io.MultiWriter(buf, sum)
@@ -71,7 +73,6 @@ func (w *wrapperMtprotoFrame) Read() (conntypes.Packet, error) { // nolint: funl
} }
buf.Reset() buf.Reset()
buf.Grow(int(messageLength) - 4 - 4)
if _, err := io.CopyN(writer, w.parent, int64(messageLength)-4-4); err != nil { if _, err := io.CopyN(writer, w.parent, int64(messageLength)-4-4); err != nil {
return nil, fmt.Errorf("cannot read the message frame: %w", err) return nil, fmt.Errorf("cannot read the message frame: %w", err)
@@ -113,8 +114,8 @@ func (w *wrapperMtprotoFrame) Write(p conntypes.Packet) error {
messageLength := 4 + 4 + len(p) + 4 messageLength := 4 + 4 + len(p) + 4
paddingLength := (aes.BlockSize - messageLength%aes.BlockSize) % aes.BlockSize paddingLength := (aes.BlockSize - messageLength%aes.BlockSize) % aes.BlockSize
buf := &bytes.Buffer{} buf := acquireMtprotoFrameBytesBuffer()
buf.Grow(messageLength + paddingLength) defer releaseMtprotoFrameBytesBuffer(buf)
binary.Write(buf, binary.LittleEndian, uint32(messageLength)) // nolint: errcheck binary.Write(buf, binary.LittleEndian, uint32(messageLength)) // nolint: errcheck
binary.Write(buf, binary.LittleEndian, w.writeSeqNo) // nolint: errcheck binary.Write(buf, binary.LittleEndian, w.writeSeqNo) // nolint: errcheck
+23
View File
@@ -0,0 +1,23 @@
package packet
import (
"bytes"
"sync"
)
var (
poolMtprotoFrameBytesBuffer = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func acquireMtprotoFrameBytesBuffer() *bytes.Buffer {
return poolMtprotoFrameBytesBuffer.Get().(*bytes.Buffer)
}
func releaseMtprotoFrameBytesBuffer(buf *bytes.Buffer) {
buf.Reset()
poolMtprotoFrameBytesBuffer.Put(buf)
}
+3 -1
View File
@@ -88,7 +88,9 @@ func (w *wrapperClientAbridged) Write(packet conntypes.Packet, acks *conntypes.C
return nil return nil
case packetLength < clientAbridgedLargePacketLength: case packetLength < clientAbridgedLargePacketLength:
length24 := utils.ToUint24(uint32(packetLength)) length24 := utils.ToUint24(uint32(packetLength))
buf := bytes.Buffer{}
buf := acquireClientBytesBuffer()
defer releaseClientBytesBuffer(buf)
buf.WriteByte(byte(clientAbridgedSmallPacketLength)) buf.WriteByte(byte(clientAbridgedSmallPacketLength))
buf.Write(length24[:]) buf.Write(length24[:])
@@ -1,7 +1,6 @@
package packetack package packetack
import ( import (
"bytes"
"encoding/binary" "encoding/binary"
"fmt" "fmt"
"math/rand" "math/rand"
@@ -35,11 +34,13 @@ func (w *wrapperClientIntermediateSecure) Write(packet conntypes.Packet, acks *c
return nil return nil
} }
buf := bytes.Buffer{} buf := acquireClientBytesBuffer()
defer releaseClientBytesBuffer(buf)
paddingLength := rand.Intn(4) paddingLength := rand.Intn(4)
buf.Grow(4 + len(packet) + paddingLength) buf.Grow(4 + len(packet) + paddingLength)
binary.Write(&buf, binary.LittleEndian, uint32(len(packet)+paddingLength)) // nolint: errcheck binary.Write(buf, binary.LittleEndian, uint32(len(packet)+paddingLength)) // nolint: errcheck
buf.Write(packet) buf.Write(packet)
buf.Write(make([]byte, paddingLength)) buf.Write(make([]byte, paddingLength))
+23
View File
@@ -0,0 +1,23 @@
package packetack
import (
"bytes"
"sync"
)
var (
poolClientBytesBuffer = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func acquireClientBytesBuffer() *bytes.Buffer {
return poolClientBytesBuffer.Get().(*bytes.Buffer)
}
func releaseClientBytesBuffer(buf *bytes.Buffer) {
buf.Reset()
poolClientBytesBuffer.Put(buf)
}
+2 -1
View File
@@ -23,8 +23,8 @@ type wrapperProxy struct {
func (w *wrapperProxy) Write(packet conntypes.Packet, acks *conntypes.ConnectionAcks) error { func (w *wrapperProxy) Write(packet conntypes.Packet, acks *conntypes.ConnectionAcks) error {
buf := bytes.Buffer{} buf := bytes.Buffer{}
flags := w.flags flags := w.flags
if acks.Quick { if acks.Quick {
flags |= rpc.ProxyRequestFlagsQuickAck flags |= rpc.ProxyRequestFlagsQuickAck
} }
@@ -43,6 +43,7 @@ func (w *wrapperProxy) Write(packet conntypes.Packet, acks *conntypes.Connection
buf.WriteByte(byte(len(config.C.AdTag))) buf.WriteByte(byte(len(config.C.AdTag)))
buf.Write(config.C.AdTag) buf.Write(config.C.AdTag)
buf.Write(make([]byte, (4-buf.Len()%4)%4)) buf.Write(make([]byte, (4-buf.Len()%4)%4))
buf.Grow(len(packet))
buf.Write(packet) buf.Write(packet)
return w.proxy.Write(buf.Bytes()) return w.proxy.Write(buf.Bytes())
+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)
}