Correctly manage partial writes

This commit is contained in:
9seconds
2021-03-29 10:26:52 +03:00
parent db8614999a
commit 3992054560
7 changed files with 75 additions and 38 deletions
+25 -6
View File
@@ -1,12 +1,22 @@
package record package record
import "sync" import (
"bytes"
"sync"
)
var recordPool = sync.Pool{ var (
New: func() interface{} { recordPool = sync.Pool{
return &Record{} New: func() interface{} {
}, return &Record{}
} },
}
bytesBufferPool = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func AcquireRecord() *Record { func AcquireRecord() *Record {
return recordPool.Get().(*Record) return recordPool.Get().(*Record)
@@ -16,3 +26,12 @@ func ReleaseRecord(r *Record) {
r.Reset() r.Reset()
recordPool.Put(r) recordPool.Put(r)
} }
func acquireBytesBuffer() *bytes.Buffer {
return bytesBufferPool.Get().(*bytes.Buffer)
}
func releaseBytesBuffer(buf *bytes.Buffer) {
buf.Reset()
bytesBufferPool.Put(buf)
}
+11 -15
View File
@@ -61,26 +61,22 @@ func (r *Record) Read(reader io.Reader) error {
} }
func (r *Record) Dump(writer io.Writer) error { func (r *Record) Dump(writer io.Writer) error {
buf := [2]byte{byte(r.Type), 0} buf := acquireBytesBuffer()
defer releaseBytesBuffer(buf)
if _, err := writer.Write(buf[:1]); err != nil { bufSlice := [2]byte{byte(r.Type), 0}
return fmt.Errorf("cannot dump type: %w", err) buf.Write(bufSlice[:1])
}
binary.BigEndian.PutUint16(buf[:], uint16(r.Version)) binary.BigEndian.PutUint16(bufSlice[:], uint16(r.Version))
buf.Write(bufSlice[:])
if _, err := writer.Write(buf[:]); err != nil { binary.BigEndian.PutUint16(bufSlice[:], uint16(r.Payload.Len()))
return fmt.Errorf("cannot dump version: %w", err) buf.Write(bufSlice[:])
}
binary.BigEndian.PutUint16(buf[:], uint16(r.Payload.Len())) buf.Write(r.Payload.Bytes())
if _, err := writer.Write(buf[:]); err != nil { if _, err := buf.WriteTo(writer); err != nil {
return fmt.Errorf("cannot dump payload length: %w", err) return fmt.Errorf("cannot dump record: %w", err)
}
if _, err := writer.Write(r.Payload.Bytes()); err != nil {
return fmt.Errorf("cannot dump payload: %w", err)
} }
return nil return nil
@@ -58,7 +58,7 @@ func (suite *ClientHandshakeTestSuite) TestOk() {
copy(writeData, arr) copy(writeData, arr)
}) })
conn := &obfuscated2.Conn{ conn := obfuscated2.Conn{
Conn: connMock, Conn: connMock,
Encryptor: encryptor, Encryptor: encryptor,
Decryptor: decryptor, Decryptor: decryptor,
+12 -7
View File
@@ -10,11 +10,9 @@ type Conn struct {
Encryptor cipher.Stream Encryptor cipher.Stream
Decryptor cipher.Stream Decryptor cipher.Stream
writeBuf []byte
} }
func (c *Conn) Read(p []byte) (int, error) { func (c Conn) Read(p []byte) (int, error) {
n, err := c.Conn.Read(p) n, err := c.Conn.Read(p)
if err != nil { if err != nil {
return n, err // nolint: wrapcheck return n, err // nolint: wrapcheck
@@ -25,9 +23,16 @@ func (c *Conn) Read(p []byte) (int, error) {
return n, nil return n, nil
} }
func (c *Conn) Write(p []byte) (int, error) { func (c Conn) Write(p []byte) (int, error) {
c.writeBuf = append(c.writeBuf[:0], p...) buf := acquireBytesBuffer()
c.Encryptor.XORKeyStream(c.writeBuf, p) defer releaseBytesBuffer(buf)
return c.Conn.Write(c.writeBuf) buf.Write(p)
payload := buf.Bytes()
c.Encryptor.XORKeyStream(payload, payload)
n, err := buf.WriteTo(c.Conn)
return int(n), err // nolint: wrapcheck
} }
+22 -5
View File
@@ -1,16 +1,24 @@
package obfuscated2 package obfuscated2
import ( import (
"bytes"
"crypto/sha256" "crypto/sha256"
"hash" "hash"
"sync" "sync"
) )
var sha256HasherPool = sync.Pool{ var (
New: func() interface{} { sha256HasherPool = sync.Pool{
return sha256.New() New: func() interface{} {
}, return sha256.New()
} },
}
bytesBufferPool = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func acquireSha256Hasher() hash.Hash { func acquireSha256Hasher() hash.Hash {
return sha256HasherPool.Get().(hash.Hash) return sha256HasherPool.Get().(hash.Hash)
@@ -20,3 +28,12 @@ func releaseSha256Hasher(h hash.Hash) {
h.Reset() h.Reset()
sha256HasherPool.Put(h) sha256HasherPool.Put(h)
} }
func acquireBytesBuffer() *bytes.Buffer {
return bytesBufferPool.Get().(*bytes.Buffer)
}
func releaseBytesBuffer(buf *bytes.Buffer) {
buf.Reset()
bytesBufferPool.Put(buf)
}
@@ -17,7 +17,7 @@ type ServerHandshakeTestSuite struct {
suite.Suite suite.Suite
connMock *testlib.NetConnMock connMock *testlib.NetConnMock
proxyConn *obfuscated2.Conn proxyConn obfuscated2.Conn
encryptor cipher.Stream encryptor cipher.Stream
decryptor cipher.Stream decryptor cipher.Stream
} }
@@ -29,7 +29,7 @@ func (suite *ServerHandshakeTestSuite) SetupTest() {
encryptor, decryptor, err := obfuscated2.ServerHandshake(buf) encryptor, decryptor, err := obfuscated2.ServerHandshake(buf)
suite.NoError(err) suite.NoError(err)
suite.proxyConn = &obfuscated2.Conn{ suite.proxyConn = obfuscated2.Conn{
Conn: suite.connMock, Conn: suite.connMock,
Encryptor: encryptor, Encryptor: encryptor,
Decryptor: decryptor, Decryptor: decryptor,
+2 -2
View File
@@ -192,7 +192,7 @@ func (p *Proxy) doObfuscated2Handshake(ctx *streamContext) error {
ctx.dc = dc ctx.dc = dc
ctx.logger = ctx.logger.BindInt("dc", dc) ctx.logger = ctx.logger.BindInt("dc", dc)
ctx.clientConn = &obfuscated2.Conn{ ctx.clientConn = obfuscated2.Conn{
Conn: ctx.clientConn, Conn: ctx.clientConn,
Encryptor: encryptor, Encryptor: encryptor,
Decryptor: decryptor, Decryptor: decryptor,
@@ -214,7 +214,7 @@ func (p *Proxy) doTelegramCall(ctx *streamContext) error {
return fmt.Errorf("cannot perform obfuscated2 handshake: %w", err) return fmt.Errorf("cannot perform obfuscated2 handshake: %w", err)
} }
ctx.telegramConn = &obfuscated2.Conn{ ctx.telegramConn = obfuscated2.Conn{
Conn: connTelegramTraffic{ Conn: connTelegramTraffic{
Conn: conn, Conn: conn,
connID: ctx.connID, connID: ctx.connID,