mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 14:54:01 +03:00
Correctly manage partial writes
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
+2
-2
@@ -192,7 +192,7 @@ func (p *Proxy) doObfuscated2Handshake(ctx *streamContext) error {
|
||||
|
||||
ctx.dc = dc
|
||||
ctx.logger = ctx.logger.BindInt("dc", dc)
|
||||
ctx.clientConn = &obfuscated2.Conn{
|
||||
ctx.clientConn = obfuscated2.Conn{
|
||||
Conn: ctx.clientConn,
|
||||
Encryptor: encryptor,
|
||||
Decryptor: decryptor,
|
||||
@@ -214,7 +214,7 @@ func (p *Proxy) doTelegramCall(ctx *streamContext) error {
|
||||
return fmt.Errorf("cannot perform obfuscated2 handshake: %w", err)
|
||||
}
|
||||
|
||||
ctx.telegramConn = &obfuscated2.Conn{
|
||||
ctx.telegramConn = obfuscated2.Conn{
|
||||
Conn: connTelegramTraffic{
|
||||
Conn: conn,
|
||||
connID: ctx.connID,
|
||||
|
||||
Reference in New Issue
Block a user