mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 15:34:01 +03:00
Merge pull request #140 from 9seconds/pools
Add pool support everywhere
This commit is contained in:
@@ -63,7 +63,12 @@ func (c *ClientProtocol) tlsHandshake(conn io.ReadWriter) error {
|
||||
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 {
|
||||
return fmt.Errorf("cannot parse client hello: %w", err)
|
||||
}
|
||||
|
||||
+11
-8
@@ -28,15 +28,9 @@ func cloak(one, another io.ReadWriteCloser) {
|
||||
|
||||
wg.Add(2)
|
||||
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
io.Copy(one, another) // nolint: errcheck
|
||||
}()
|
||||
go cloakPipe(one, another, wg)
|
||||
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
io.Copy(another, one) // nolint: errcheck
|
||||
}()
|
||||
go cloakPipe(another, one, wg)
|
||||
|
||||
go func() {
|
||||
wg.Wait()
|
||||
@@ -69,3 +63,12 @@ func cloak(one, another io.ReadWriteCloser) {
|
||||
|
||||
<-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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -25,7 +25,7 @@ func (c ClientHello) Digest() []byte {
|
||||
}
|
||||
|
||||
mac := hmac.New(sha256.New, config.C.Secret)
|
||||
mac.Write(rec.Bytes()) // nolint: errcheck
|
||||
rec.WriteBytes(mac)
|
||||
computedDigest := mac.Sum(nil)
|
||||
|
||||
for i := range computedDigest {
|
||||
|
||||
+10
-3
@@ -1,5 +1,7 @@
|
||||
package tlstypes
|
||||
|
||||
import "io"
|
||||
|
||||
type RecordType uint8
|
||||
|
||||
const (
|
||||
@@ -69,11 +71,16 @@ var (
|
||||
)
|
||||
|
||||
type Byter interface {
|
||||
Bytes() []byte
|
||||
WriteBytes(io.Writer)
|
||||
Len() int
|
||||
}
|
||||
|
||||
type RawBytes []byte
|
||||
|
||||
func (r RawBytes) Bytes() []byte {
|
||||
return []byte(r)
|
||||
func (r RawBytes) WriteBytes(writer io.Writer) {
|
||||
writer.Write(r) // nolint: errcheck
|
||||
}
|
||||
|
||||
func (r RawBytes) Len() int {
|
||||
return len(r)
|
||||
}
|
||||
|
||||
+16
-9
@@ -1,7 +1,7 @@
|
||||
package tlstypes
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
|
||||
"github.com/9seconds/mtg/utils"
|
||||
)
|
||||
@@ -14,24 +14,31 @@ type Handshake struct {
|
||||
Tail Byter
|
||||
}
|
||||
|
||||
func (h *Handshake) Bytes() []byte {
|
||||
buf := bytes.Buffer{}
|
||||
packetBuf := bytes.Buffer{}
|
||||
func (h *Handshake) WriteBytes(writer io.Writer) {
|
||||
packetBuf := acquireBytesBuffer()
|
||||
defer releaseBytesBuffer(packetBuf)
|
||||
|
||||
buf.WriteByte(byte(h.Type))
|
||||
writer.Write([]byte{byte(h.Type)}) // nolint: errcheck
|
||||
|
||||
packetBuf.Write(h.Version.Bytes())
|
||||
packetBuf.Write(h.Random[:])
|
||||
packetBuf.WriteByte(byte(len(h.SessionID)))
|
||||
packetBuf.Write(h.SessionID)
|
||||
packetBuf.Write(h.Tail.Bytes())
|
||||
h.Tail.WriteBytes(packetBuf)
|
||||
|
||||
sizeUint24 := utils.ToUint24(uint32(packetBuf.Len()))
|
||||
sizeUint24Bytes := sizeUint24[:]
|
||||
sizeUint24Bytes[0], sizeUint24Bytes[2] = sizeUint24Bytes[2], sizeUint24Bytes[0]
|
||||
|
||||
buf.Write(sizeUint24Bytes)
|
||||
packetBuf.WriteTo(&buf) // nolint: errcheck
|
||||
writer.Write(sizeUint24Bytes) // 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()
|
||||
}
|
||||
|
||||
@@ -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
@@ -15,16 +15,15 @@ type Record struct {
|
||||
Data Byter
|
||||
}
|
||||
|
||||
func (r Record) Bytes() []byte {
|
||||
buf := bytes.Buffer{}
|
||||
data := r.Data.Bytes()
|
||||
func (r Record) WriteBytes(writer io.Writer) {
|
||||
writer.Write([]byte{byte(r.Type)}) // nolint: errcheck
|
||||
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))
|
||||
buf.Write(r.Version.Bytes())
|
||||
binary.Write(&buf, binary.BigEndian, uint16(len(data))) // nolint: errcheck
|
||||
buf.Write(data)
|
||||
|
||||
return buf.Bytes()
|
||||
func (r Record) Len() int {
|
||||
return 1 + 2 + 2 + r.Data.Len()
|
||||
}
|
||||
|
||||
func ReadRecord(reader io.Reader) (Record, error) {
|
||||
|
||||
@@ -20,20 +20,22 @@ type ServerHello struct {
|
||||
}
|
||||
|
||||
func (s ServerHello) WelcomePacket() []byte {
|
||||
buf := &bytes.Buffer{}
|
||||
|
||||
s.Random = [32]byte{}
|
||||
rec := Record{
|
||||
Type: RecordTypeHandshake,
|
||||
Version: Version12,
|
||||
Data: &s,
|
||||
}
|
||||
buf := bytes.NewBuffer(rec.Bytes())
|
||||
rec.WriteBytes(buf)
|
||||
|
||||
recChangeCipher := Record{
|
||||
Type: RecordTypeChangeCipherSpec,
|
||||
Version: Version12,
|
||||
Data: RawBytes([]byte{0x01}),
|
||||
}
|
||||
buf.Write(recChangeCipher.Bytes())
|
||||
recChangeCipher.WriteBytes(buf)
|
||||
|
||||
hostCert := make([]byte, 1024+mrand.Intn(3092))
|
||||
rand.Read(hostCert) // nolint: errcheck
|
||||
@@ -43,7 +45,8 @@ func (s ServerHello) WelcomePacket() []byte {
|
||||
Version: Version12,
|
||||
Data: RawBytes(hostCert),
|
||||
}
|
||||
buf.Write(recData.Bytes())
|
||||
recData.WriteBytes(buf)
|
||||
|
||||
packet := buf.Bytes()
|
||||
|
||||
mac := hmac.New(sha256.New, config.C.Secret)
|
||||
|
||||
@@ -42,7 +42,9 @@ type wrapperMtprotoFrame struct {
|
||||
}
|
||||
|
||||
func (w *wrapperMtprotoFrame) Read() (conntypes.Packet, error) { // nolint: funlen
|
||||
buf := &bytes.Buffer{}
|
||||
buf := acquireMtprotoFrameBytesBuffer()
|
||||
defer releaseMtprotoFrameBytesBuffer(buf)
|
||||
|
||||
sum := crc32.NewIEEE()
|
||||
writer := io.MultiWriter(buf, sum)
|
||||
|
||||
@@ -71,7 +73,6 @@ func (w *wrapperMtprotoFrame) Read() (conntypes.Packet, error) { // nolint: funl
|
||||
}
|
||||
|
||||
buf.Reset()
|
||||
buf.Grow(int(messageLength) - 4 - 4)
|
||||
|
||||
if _, err := io.CopyN(writer, w.parent, int64(messageLength)-4-4); err != nil {
|
||||
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
|
||||
paddingLength := (aes.BlockSize - messageLength%aes.BlockSize) % aes.BlockSize
|
||||
|
||||
buf := &bytes.Buffer{}
|
||||
buf.Grow(messageLength + paddingLength)
|
||||
buf := acquireMtprotoFrameBytesBuffer()
|
||||
defer releaseMtprotoFrameBytesBuffer(buf)
|
||||
|
||||
binary.Write(buf, binary.LittleEndian, uint32(messageLength)) // nolint: errcheck
|
||||
binary.Write(buf, binary.LittleEndian, w.writeSeqNo) // nolint: errcheck
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -88,7 +88,9 @@ func (w *wrapperClientAbridged) Write(packet conntypes.Packet, acks *conntypes.C
|
||||
return nil
|
||||
case packetLength < clientAbridgedLargePacketLength:
|
||||
length24 := utils.ToUint24(uint32(packetLength))
|
||||
buf := bytes.Buffer{}
|
||||
|
||||
buf := acquireClientBytesBuffer()
|
||||
defer releaseClientBytesBuffer(buf)
|
||||
|
||||
buf.WriteByte(byte(clientAbridgedSmallPacketLength))
|
||||
buf.Write(length24[:])
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package packetack
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"math/rand"
|
||||
@@ -35,11 +34,13 @@ func (w *wrapperClientIntermediateSecure) Write(packet conntypes.Packet, acks *c
|
||||
return nil
|
||||
}
|
||||
|
||||
buf := bytes.Buffer{}
|
||||
buf := acquireClientBytesBuffer()
|
||||
defer releaseClientBytesBuffer(buf)
|
||||
|
||||
paddingLength := rand.Intn(4)
|
||||
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(make([]byte, paddingLength))
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -23,8 +23,8 @@ type wrapperProxy struct {
|
||||
|
||||
func (w *wrapperProxy) Write(packet conntypes.Packet, acks *conntypes.ConnectionAcks) error {
|
||||
buf := bytes.Buffer{}
|
||||
|
||||
flags := w.flags
|
||||
|
||||
if acks.Quick {
|
||||
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.Write(config.C.AdTag)
|
||||
buf.Write(make([]byte, (4-buf.Len()%4)%4))
|
||||
buf.Grow(len(packet))
|
||||
buf.Write(packet)
|
||||
|
||||
return w.proxy.Write(buf.Bytes())
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package stream
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"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) {
|
||||
sum := 0
|
||||
|
||||
buf := acquireBytesBuffer()
|
||||
defer releaseBytesBuffer(buf)
|
||||
|
||||
for _, v := range tlstypes.MakeRecords(p) {
|
||||
_, err := writeFunc(v.Bytes())
|
||||
buf.Reset()
|
||||
v.WriteBytes(buf)
|
||||
|
||||
_, err := writeFunc(buf.Bytes())
|
||||
if err != nil {
|
||||
return sum, err
|
||||
}
|
||||
|
||||
sum += len(v.Data.Bytes())
|
||||
sum += v.Data.Len()
|
||||
}
|
||||
|
||||
return sum, nil
|
||||
@@ -86,7 +93,10 @@ func NewFakeTLS(socket conntypes.StreamReadWriteCloser) conntypes.StreamReadWrit
|
||||
switch rec.Type {
|
||||
case tlstypes.RecordTypeChangeCipherSpec:
|
||||
case tlstypes.RecordTypeApplicationData:
|
||||
return rec.Data.Bytes(), nil
|
||||
buf := &bytes.Buffer{}
|
||||
rec.Data.WriteBytes(buf)
|
||||
|
||||
return buf.Bytes(), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported record type %v", rec.Type)
|
||||
}
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package stream
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/md5" // nolint: gosec
|
||||
@@ -54,7 +53,9 @@ func mtprotoDeriveKeys(purpose mtprotoCipherPurpose,
|
||||
resp *rpc.NonceResponse,
|
||||
client, remote *net.TCPAddr,
|
||||
secret []byte) ([]byte, []byte) {
|
||||
message := bytes.Buffer{}
|
||||
message := acquireBytesBuffer()
|
||||
defer releaseBytesBuffer(message)
|
||||
|
||||
message.Write(resp.Nonce) // nolint: gosec
|
||||
message.Write(req.Nonce) // nolint: gosec
|
||||
message.Write(req.CryptoTS) // nolint: gosec
|
||||
|
||||
@@ -1,11 +1,9 @@
|
||||
package stream
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/cipher"
|
||||
"fmt"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
@@ -13,23 +11,6 @@ import (
|
||||
"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 {
|
||||
encryptor 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) {
|
||||
buffer := poolWrapperObfuscated2WritePoolAcquire()
|
||||
defer poolWrapperObfuscated2WritePoolRelease(buffer)
|
||||
buffer := acquireBytesBuffer()
|
||||
defer releaseBytesBuffer(buffer)
|
||||
|
||||
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) {
|
||||
buffer := poolWrapperObfuscated2WritePoolAcquire()
|
||||
defer poolWrapperObfuscated2WritePoolRelease(buffer)
|
||||
buffer := acquireBytesBuffer()
|
||||
defer releaseBytesBuffer(buffer)
|
||||
|
||||
buffer.Write(p)
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user