mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 23:34:02 +03:00
Correctly manage partial writes
This commit is contained in:
@@ -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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user