mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-09-01 03:44:01 +03:00
Add buffered reader
This commit is contained in:
@@ -29,10 +29,9 @@ type RPCProxyRequest struct {
|
|||||||
LocalIPPort [rpcProxyRequestIPPortLength]byte
|
LocalIPPort [rpcProxyRequestIPPortLength]byte
|
||||||
ADTag []byte
|
ADTag []byte
|
||||||
Extras Extras
|
Extras Extras
|
||||||
Message *bytes.Buffer
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *RPCProxyRequest) Bytes() []byte {
|
func (r *RPCProxyRequest) Bytes(message []byte) []byte {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
|
|
||||||
flags := r.Flags
|
flags := r.Flags
|
||||||
@@ -40,8 +39,7 @@ func (r *RPCProxyRequest) Bytes() []byte {
|
|||||||
flags |= RPCProxyRequestFlagsQuickAck
|
flags |= RPCProxyRequestFlagsQuickAck
|
||||||
}
|
}
|
||||||
|
|
||||||
messageBytes := r.Message.Bytes()
|
if bytes.HasPrefix(message, rpcProxyRequestFlagsEncryptedPrefix[:]) {
|
||||||
if bytes.HasPrefix(messageBytes, rpcProxyRequestFlagsEncryptedPrefix[:]) {
|
|
||||||
flags |= RPCProxyRequestFlagsEncrypted
|
flags |= RPCProxyRequestFlagsEncrypted
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -58,9 +56,7 @@ func (r *RPCProxyRequest) Bytes() []byte {
|
|||||||
for i := 0; i < (buf.Len() % 4); i++ {
|
for i := 0; i < (buf.Len() % 4); i++ {
|
||||||
buf.WriteByte(0x00)
|
buf.WriteByte(0x00)
|
||||||
}
|
}
|
||||||
if r.Message != nil {
|
buf.Write(message)
|
||||||
buf.Write(messageBytes)
|
|
||||||
}
|
|
||||||
|
|
||||||
return buf.Bytes()
|
return buf.Bytes()
|
||||||
}
|
}
|
||||||
|
|||||||
+42
-61
@@ -22,11 +22,11 @@ const (
|
|||||||
var frameRWCPadding = [4]byte{0x04, 0x00, 0x00, 0x00}
|
var frameRWCPadding = [4]byte{0x04, 0x00, 0x00, 0x00}
|
||||||
|
|
||||||
type FrameRWC struct {
|
type FrameRWC struct {
|
||||||
conn wrappers.ReadWriteCloserWithAddr
|
wrappers.BufferedReader
|
||||||
|
|
||||||
|
conn wrappers.ReadWriteCloserWithAddr
|
||||||
readSeqNo int32
|
readSeqNo int32
|
||||||
writeSeqNo int32
|
writeSeqNo int32
|
||||||
readBuf *bytes.Buffer
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *FrameRWC) Write(buf []byte) (int, error) {
|
func (f *FrameRWC) Write(buf []byte) (int, error) {
|
||||||
@@ -54,51 +54,49 @@ func (f *FrameRWC) Write(buf []byte) (int, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *FrameRWC) Read(p []byte) (int, error) {
|
func (f *FrameRWC) Read(p []byte) (int, error) {
|
||||||
if f.readBuf.Len() > 0 {
|
return f.BufferedRead(p, func() error {
|
||||||
return f.flush(p)
|
buf := &bytes.Buffer{}
|
||||||
}
|
for {
|
||||||
|
buf.Reset()
|
||||||
|
if _, err := io.CopyN(buf, f.conn, 4); err != nil {
|
||||||
|
return errors.Annotate(err, "Cannot read frame padding")
|
||||||
|
}
|
||||||
|
if !bytes.Equal(buf.Bytes(), frameRWCPadding[:]) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
messageLength := binary.LittleEndian.Uint32(buf.Bytes())
|
||||||
|
if messageLength%4 != 0 || messageLength < frameRWCMinMessageLength || messageLength > frameRWCMaxMessageLength {
|
||||||
|
return errors.Errorf("Incorrect frame message length %d", messageLength)
|
||||||
|
}
|
||||||
|
sum := crc32.NewIEEE()
|
||||||
|
sum.Write(buf.Bytes())
|
||||||
|
|
||||||
buf := &bytes.Buffer{}
|
|
||||||
for {
|
|
||||||
buf.Reset()
|
buf.Reset()
|
||||||
if _, err := io.CopyN(buf, f.conn, 4); err != nil {
|
buf.Grow(int(messageLength) - 4) // -4 because we already read the first number
|
||||||
return 0, errors.Annotate(err, "Cannot read frame padding")
|
if _, err := io.CopyN(buf, f.conn, int64(messageLength)-4); err != nil {
|
||||||
|
return errors.Annotate(err, "Cannot read the message frame")
|
||||||
}
|
}
|
||||||
if !bytes.Equal(buf.Bytes(), frameRWCPadding[:]) {
|
sum.Write(buf.Bytes())
|
||||||
break
|
|
||||||
|
var seqNo int32
|
||||||
|
binary.Read(buf, binary.LittleEndian, seqNo)
|
||||||
|
if seqNo != f.readSeqNo {
|
||||||
|
return errors.Errorf("Unexpected sequence number %d (wait for %d)", seqNo, f.readSeqNo)
|
||||||
}
|
}
|
||||||
}
|
f.readSeqNo++
|
||||||
|
|
||||||
messageLength := binary.LittleEndian.Uint32(buf.Bytes())
|
data := buf.Bytes()[:int(messageLength)-4-4-4]
|
||||||
if messageLength%4 != 0 || messageLength < frameRWCMinMessageLength || messageLength > frameRWCMaxMessageLength {
|
checksum := binary.LittleEndian.Uint32(buf.Bytes()[int(messageLength)-4-4-4:])
|
||||||
return 0, errors.Errorf("Incorrect frame message length %d", messageLength)
|
if checksum != sum.Sum32() {
|
||||||
}
|
return errors.Errorf("CRC32 checksum mismatch. Wait for %d, got %d", sum.Sum32(), checksum)
|
||||||
sum := crc32.NewIEEE()
|
|
||||||
sum.Write(buf.Bytes())
|
|
||||||
|
|
||||||
buf.Reset()
|
}
|
||||||
buf.Grow(int(messageLength) - 4) // -4 because we already read the first number
|
f.Buffer.Write(data)
|
||||||
if _, err := io.CopyN(buf, f.conn, int64(messageLength)-4); err != nil {
|
|
||||||
return 0, errors.Annotate(err, "Cannot read the message frame")
|
|
||||||
}
|
|
||||||
sum.Write(buf.Bytes())
|
|
||||||
|
|
||||||
var seqNo int32
|
return nil
|
||||||
binary.Read(buf, binary.LittleEndian, seqNo)
|
})
|
||||||
if seqNo != f.readSeqNo {
|
|
||||||
return 0, errors.Errorf("Unexpected sequence number %d (wait for %d)", seqNo, f.readSeqNo)
|
|
||||||
}
|
|
||||||
f.readSeqNo++
|
|
||||||
|
|
||||||
data := buf.Bytes()[:int(messageLength)-4-4-4]
|
|
||||||
checksum := binary.LittleEndian.Uint32(buf.Bytes()[int(messageLength)-4-4-4:])
|
|
||||||
if checksum != sum.Sum32() {
|
|
||||||
return 0, errors.Errorf("CRC32 checksum mismatch. Wait for %d, got %d", sum.Sum32(), checksum)
|
|
||||||
|
|
||||||
}
|
|
||||||
f.readBuf.Write(data)
|
|
||||||
|
|
||||||
return f.flush(p)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *FrameRWC) Close() error {
|
func (f *FrameRWC) Close() error {
|
||||||
@@ -109,28 +107,11 @@ func (f *FrameRWC) Addr() *net.TCPAddr {
|
|||||||
return f.conn.Addr()
|
return f.conn.Addr()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *FrameRWC) flush(p []byte) (int, error) {
|
|
||||||
sizeToRead := len(p)
|
|
||||||
if f.readBuf.Len() < sizeToRead {
|
|
||||||
sizeToRead = f.readBuf.Len()
|
|
||||||
}
|
|
||||||
|
|
||||||
data := f.readBuf.Bytes()
|
|
||||||
copy(p, data[:sizeToRead])
|
|
||||||
if sizeToRead == f.readBuf.Len() {
|
|
||||||
f.readBuf.Reset()
|
|
||||||
} else {
|
|
||||||
f.readBuf = bytes.NewBuffer(data[sizeToRead:])
|
|
||||||
}
|
|
||||||
|
|
||||||
return sizeToRead, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewFrameRWC(conn wrappers.ReadWriteCloserWithAddr, seqNo int32) wrappers.ReadWriteCloserWithAddr {
|
func NewFrameRWC(conn wrappers.ReadWriteCloserWithAddr, seqNo int32) wrappers.ReadWriteCloserWithAddr {
|
||||||
return &FrameRWC{
|
return &FrameRWC{
|
||||||
conn: conn,
|
BufferedReader: wrappers.NewBufferedReader(),
|
||||||
readSeqNo: seqNo,
|
conn: conn,
|
||||||
writeSeqNo: seqNo,
|
readSeqNo: seqNo,
|
||||||
readBuf: &bytes.Buffer{},
|
writeSeqNo: seqNo,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+17
-46
@@ -1,7 +1,6 @@
|
|||||||
package wrappers
|
package wrappers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"crypto/aes"
|
"crypto/aes"
|
||||||
"crypto/cipher"
|
"crypto/cipher"
|
||||||
"net"
|
"net"
|
||||||
@@ -10,7 +9,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type BlockCipherReadWriteCloserWithAddr struct {
|
type BlockCipherReadWriteCloserWithAddr struct {
|
||||||
buf *bytes.Buffer
|
BufferedReader
|
||||||
|
|
||||||
conn ReadWriteCloserWithAddr
|
conn ReadWriteCloserWithAddr
|
||||||
encryptor cipher.BlockMode
|
encryptor cipher.BlockMode
|
||||||
@@ -18,20 +17,19 @@ type BlockCipherReadWriteCloserWithAddr struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *BlockCipherReadWriteCloserWithAddr) Read(p []byte) (int, error) {
|
func (c *BlockCipherReadWriteCloserWithAddr) Read(p []byte) (int, error) {
|
||||||
if c.buf.Len() > 0 {
|
return c.BufferedRead(p, func() error {
|
||||||
return c.flush(p)
|
bufferLength := c.Buffer.Len()
|
||||||
}
|
for bufferLength%aes.BlockSize != 0 || bufferLength == 0 {
|
||||||
|
n, err := c.conn.Read(p)
|
||||||
for c.buf.Len() == 0 || c.buf.Len()%aes.BlockSize != 0 {
|
if err != nil {
|
||||||
n, err := c.conn.Read(p)
|
return errors.Annotate(err, "Cannot read from socket")
|
||||||
if err != nil {
|
}
|
||||||
return 0, errors.Annotate(err, "Cannot read from socket")
|
c.Buffer.Write(p[:n])
|
||||||
}
|
}
|
||||||
c.buf.Write(p[:n])
|
c.decryptor.CryptBlocks(c.Buffer.Bytes(), c.Buffer.Bytes())
|
||||||
}
|
|
||||||
c.decryptor.CryptBlocks(c.buf.Bytes(), c.buf.Bytes())
|
|
||||||
|
|
||||||
return c.flush(p)
|
return nil
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *BlockCipherReadWriteCloserWithAddr) Write(p []byte) (int, error) {
|
func (c *BlockCipherReadWriteCloserWithAddr) Write(p []byte) (int, error) {
|
||||||
@@ -39,19 +37,13 @@ func (c *BlockCipherReadWriteCloserWithAddr) Write(p []byte) (int, error) {
|
|||||||
return 0, errors.Errorf("Incorrect block size %d", len(p))
|
return 0, errors.Errorf("Incorrect block size %d", len(p))
|
||||||
}
|
}
|
||||||
|
|
||||||
buf := getBuffer()
|
encrypted := make([]byte, len(p))
|
||||||
defer putBuffer(buf)
|
|
||||||
buf.Grow(len(p))
|
|
||||||
buf.Write(p)
|
|
||||||
|
|
||||||
encrypted := buf.Bytes()
|
|
||||||
c.encryptor.CryptBlocks(encrypted, p)
|
c.encryptor.CryptBlocks(encrypted, p)
|
||||||
|
|
||||||
return c.conn.Write(encrypted)
|
return c.conn.Write(encrypted)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *BlockCipherReadWriteCloserWithAddr) Close() error {
|
func (c *BlockCipherReadWriteCloserWithAddr) Close() error {
|
||||||
defer putBuffer(c.buf)
|
|
||||||
return c.conn.Close()
|
return c.conn.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -59,32 +51,11 @@ func (c *BlockCipherReadWriteCloserWithAddr) Addr() *net.TCPAddr {
|
|||||||
return c.conn.Addr()
|
return c.conn.Addr()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *BlockCipherReadWriteCloserWithAddr) flush(p []byte) (int, error) {
|
|
||||||
sizeToRead := len(p)
|
|
||||||
if c.buf.Len() < sizeToRead {
|
|
||||||
sizeToRead = c.buf.Len()
|
|
||||||
}
|
|
||||||
|
|
||||||
data := c.buf.Bytes()
|
|
||||||
copy(p, data[:sizeToRead])
|
|
||||||
if sizeToRead == c.buf.Len() {
|
|
||||||
c.buf.Reset()
|
|
||||||
} else {
|
|
||||||
newBuf := getBuffer()
|
|
||||||
newBuf.Write(data[sizeToRead:])
|
|
||||||
|
|
||||||
putBuffer(c.buf)
|
|
||||||
c.buf = newBuf
|
|
||||||
}
|
|
||||||
|
|
||||||
return sizeToRead, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewBlockCipherRWC(conn ReadWriteCloserWithAddr, encryptor, decryptor cipher.BlockMode) ReadWriteCloserWithAddr {
|
func NewBlockCipherRWC(conn ReadWriteCloserWithAddr, encryptor, decryptor cipher.BlockMode) ReadWriteCloserWithAddr {
|
||||||
return &BlockCipherReadWriteCloserWithAddr{
|
return &BlockCipherReadWriteCloserWithAddr{
|
||||||
buf: getBuffer(),
|
BufferedReader: NewBufferedReader(),
|
||||||
conn: conn,
|
conn: conn,
|
||||||
encryptor: encryptor,
|
encryptor: encryptor,
|
||||||
decryptor: decryptor,
|
decryptor: decryptor,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
package wrappers
|
||||||
|
|
||||||
|
import "bytes"
|
||||||
|
|
||||||
|
type BufferedReader struct {
|
||||||
|
Buffer *bytes.Buffer
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BufferedReader) BufferedRead(p []byte, callback func() error) (int, error) {
|
||||||
|
if b.Buffer.Len() > 0 {
|
||||||
|
return b.flush(p)
|
||||||
|
}
|
||||||
|
if err := callback(); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return b.flush(p)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BufferedReader) flush(p []byte) (int, error) {
|
||||||
|
sizeToRead := len(p)
|
||||||
|
if b.Buffer.Len() < sizeToRead {
|
||||||
|
sizeToRead = b.Buffer.Len()
|
||||||
|
}
|
||||||
|
|
||||||
|
data := b.Buffer.Bytes()
|
||||||
|
copy(p, data[:sizeToRead])
|
||||||
|
if sizeToRead == b.Buffer.Len() {
|
||||||
|
b.Buffer.Reset()
|
||||||
|
} else {
|
||||||
|
b.Buffer = bytes.NewBuffer(data[sizeToRead:])
|
||||||
|
}
|
||||||
|
|
||||||
|
return sizeToRead, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewBufferedReader() BufferedReader {
|
||||||
|
return BufferedReader{Buffer: &bytes.Buffer{}}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user