Add buffered reader

This commit is contained in:
9seconds
2018-07-03 10:10:00 +03:00
parent 5f54a235d2
commit 586ffab01b
4 changed files with 100 additions and 114 deletions
+3 -7
View File
@@ -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()
} }
+12 -31
View File
@@ -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,15 +54,12 @@ 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{} buf := &bytes.Buffer{}
for { for {
buf.Reset() buf.Reset()
if _, err := io.CopyN(buf, f.conn, 4); err != nil { if _, err := io.CopyN(buf, f.conn, 4); err != nil {
return 0, errors.Annotate(err, "Cannot read frame padding") return errors.Annotate(err, "Cannot read frame padding")
} }
if !bytes.Equal(buf.Bytes(), frameRWCPadding[:]) { if !bytes.Equal(buf.Bytes(), frameRWCPadding[:]) {
break break
@@ -71,7 +68,7 @@ func (f *FrameRWC) Read(p []byte) (int, error) {
messageLength := binary.LittleEndian.Uint32(buf.Bytes()) messageLength := binary.LittleEndian.Uint32(buf.Bytes())
if messageLength%4 != 0 || messageLength < frameRWCMinMessageLength || messageLength > frameRWCMaxMessageLength { if messageLength%4 != 0 || messageLength < frameRWCMinMessageLength || messageLength > frameRWCMaxMessageLength {
return 0, errors.Errorf("Incorrect frame message length %d", messageLength) return errors.Errorf("Incorrect frame message length %d", messageLength)
} }
sum := crc32.NewIEEE() sum := crc32.NewIEEE()
sum.Write(buf.Bytes()) sum.Write(buf.Bytes())
@@ -79,26 +76,27 @@ func (f *FrameRWC) Read(p []byte) (int, error) {
buf.Reset() buf.Reset()
buf.Grow(int(messageLength) - 4) // -4 because we already read the first number buf.Grow(int(messageLength) - 4) // -4 because we already read the first number
if _, err := io.CopyN(buf, f.conn, int64(messageLength)-4); err != nil { if _, err := io.CopyN(buf, f.conn, int64(messageLength)-4); err != nil {
return 0, errors.Annotate(err, "Cannot read the message frame") return errors.Annotate(err, "Cannot read the message frame")
} }
sum.Write(buf.Bytes()) sum.Write(buf.Bytes())
var seqNo int32 var seqNo int32
binary.Read(buf, binary.LittleEndian, seqNo) binary.Read(buf, binary.LittleEndian, seqNo)
if seqNo != f.readSeqNo { if seqNo != f.readSeqNo {
return 0, errors.Errorf("Unexpected sequence number %d (wait for %d)", seqNo, f.readSeqNo) return errors.Errorf("Unexpected sequence number %d (wait for %d)", seqNo, f.readSeqNo)
} }
f.readSeqNo++ f.readSeqNo++
data := buf.Bytes()[:int(messageLength)-4-4-4] data := buf.Bytes()[:int(messageLength)-4-4-4]
checksum := binary.LittleEndian.Uint32(buf.Bytes()[int(messageLength)-4-4-4:]) checksum := binary.LittleEndian.Uint32(buf.Bytes()[int(messageLength)-4-4-4:])
if checksum != sum.Sum32() { if checksum != sum.Sum32() {
return 0, errors.Errorf("CRC32 checksum mismatch. Wait for %d, got %d", sum.Sum32(), checksum) return errors.Errorf("CRC32 checksum mismatch. Wait for %d, got %d", sum.Sum32(), checksum)
} }
f.readBuf.Write(data) f.Buffer.Write(data)
return f.flush(p) return nil
})
} }
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{
BufferedReader: wrappers.NewBufferedReader(),
conn: conn, conn: conn,
readSeqNo: seqNo, readSeqNo: seqNo,
writeSeqNo: seqNo, writeSeqNo: seqNo,
readBuf: &bytes.Buffer{},
} }
} }
+11 -40
View File
@@ -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 {
for c.buf.Len() == 0 || c.buf.Len()%aes.BlockSize != 0 {
n, err := c.conn.Read(p) n, err := c.conn.Read(p)
if err != nil { if err != nil {
return 0, errors.Annotate(err, "Cannot read from socket") return errors.Annotate(err, "Cannot read from socket")
} }
c.buf.Write(p[:n]) c.Buffer.Write(p[:n])
} }
c.decryptor.CryptBlocks(c.buf.Bytes(), c.buf.Bytes()) c.decryptor.CryptBlocks(c.Buffer.Bytes(), c.Buffer.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,30 +51,9 @@ 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,
+38
View File
@@ -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{}}
}