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
import "sync"
import (
"bytes"
"sync"
)
var recordPool = sync.Pool{
New: func() interface{} {
return &Record{}
},
}
var (
recordPool = sync.Pool{
New: func() interface{} {
return &Record{}
},
}
bytesBufferPool = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func AcquireRecord() *Record {
return recordPool.Get().(*Record)
@@ -16,3 +26,12 @@ func ReleaseRecord(r *Record) {
r.Reset()
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 {
buf := [2]byte{byte(r.Type), 0}
buf := acquireBytesBuffer()
defer releaseBytesBuffer(buf)
if _, err := writer.Write(buf[:1]); err != nil {
return fmt.Errorf("cannot dump type: %w", err)
}
bufSlice := [2]byte{byte(r.Type), 0}
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 {
return fmt.Errorf("cannot dump version: %w", err)
}
binary.BigEndian.PutUint16(bufSlice[:], uint16(r.Payload.Len()))
buf.Write(bufSlice[:])
binary.BigEndian.PutUint16(buf[:], uint16(r.Payload.Len()))
buf.Write(r.Payload.Bytes())
if _, err := writer.Write(buf[:]); err != nil {
return fmt.Errorf("cannot dump payload length: %w", err)
}
if _, err := writer.Write(r.Payload.Bytes()); err != nil {
return fmt.Errorf("cannot dump payload: %w", err)
if _, err := buf.WriteTo(writer); err != nil {
return fmt.Errorf("cannot dump record: %w", err)
}
return nil
@@ -58,7 +58,7 @@ func (suite *ClientHandshakeTestSuite) TestOk() {
copy(writeData, arr)
})
conn := &obfuscated2.Conn{
conn := obfuscated2.Conn{
Conn: connMock,
Encryptor: encryptor,
Decryptor: decryptor,
+12 -7
View File
@@ -10,11 +10,9 @@ type Conn struct {
Encryptor 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)
if err != nil {
return n, err // nolint: wrapcheck
@@ -25,9 +23,16 @@ func (c *Conn) Read(p []byte) (int, error) {
return n, nil
}
func (c *Conn) Write(p []byte) (int, error) {
c.writeBuf = append(c.writeBuf[:0], p...)
c.Encryptor.XORKeyStream(c.writeBuf, p)
func (c Conn) Write(p []byte) (int, error) {
buf := acquireBytesBuffer()
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
import (
"bytes"
"crypto/sha256"
"hash"
"sync"
)
var sha256HasherPool = sync.Pool{
New: func() interface{} {
return sha256.New()
},
}
var (
sha256HasherPool = sync.Pool{
New: func() interface{} {
return sha256.New()
},
}
bytesBufferPool = sync.Pool{
New: func() interface{} {
return &bytes.Buffer{}
},
}
)
func acquireSha256Hasher() hash.Hash {
return sha256HasherPool.Get().(hash.Hash)
@@ -20,3 +28,12 @@ func releaseSha256Hasher(h hash.Hash) {
h.Reset()
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
connMock *testlib.NetConnMock
proxyConn *obfuscated2.Conn
proxyConn obfuscated2.Conn
encryptor cipher.Stream
decryptor cipher.Stream
}
@@ -29,7 +29,7 @@ func (suite *ServerHandshakeTestSuite) SetupTest() {
encryptor, decryptor, err := obfuscated2.ServerHandshake(buf)
suite.NoError(err)
suite.proxyConn = &obfuscated2.Conn{
suite.proxyConn = obfuscated2.Conn{
Conn: suite.connMock,
Encryptor: encryptor,
Decryptor: decryptor,