Fix infinite hangs

This commit is contained in:
9seconds
2018-07-05 20:46:44 +03:00
parent 53ef81a4fc
commit 5413a5b090
5 changed files with 18 additions and 15 deletions
+3 -3
View File
@@ -28,6 +28,9 @@ type AbridgedReadWriteCloserWithAddr struct {
func (a *AbridgedReadWriteCloserWithAddr) Read(p []byte) (int, error) { func (a *AbridgedReadWriteCloserWithAddr) Read(p []byte) (int, error) {
return a.BufferedRead(p, func() error { return a.BufferedRead(p, func() error {
a.opts.QuickAck = false
a.opts.SimpleAck = false
buf := &bytes.Buffer{} buf := &bytes.Buffer{}
buf.Grow(3) buf.Grow(3)
@@ -37,7 +40,6 @@ func (a *AbridgedReadWriteCloserWithAddr) Read(p []byte) (int, error) {
msgLength := uint8(buf.Bytes()[0]) msgLength := uint8(buf.Bytes()[0])
buf.Reset() buf.Reset()
a.opts.QuickAck = false
if msgLength >= abridgedQuickAckLength { if msgLength >= abridgedQuickAckLength {
a.opts.QuickAck = true a.opts.QuickAck = true
msgLength -= 0x80 msgLength -= 0x80
@@ -78,13 +80,11 @@ func (a *AbridgedReadWriteCloserWithAddr) Write(p []byte) (int, error) {
case packetLength < abridgedLargePacketLength: case packetLength < abridgedLargePacketLength:
length24 := toUint24(uint32(packetLength)) length24 := toUint24(uint32(packetLength))
buf := &bytes.Buffer{} buf := &bytes.Buffer{}
buf.Grow(1 + 3 + len(p)) buf.Grow(1 + 3 + len(p))
buf.WriteByte(byte(abridgedSmallPacketLength)) buf.WriteByte(byte(abridgedSmallPacketLength))
buf.Write(length24[:]) buf.Write(length24[:])
buf.Write(p) buf.Write(p)
return a.conn.Write(buf.Bytes()) return a.conn.Write(buf.Bytes())
default: default:
+1 -1
View File
@@ -66,7 +66,7 @@ func (f *FrameRWC) Read(p []byte) (int, error) {
if _, err := io.CopyN(writer, f.conn, 4); err != nil { if _, err := io.CopyN(writer, f.conn, 4); err != nil {
return 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
} }
} }
+3
View File
@@ -23,6 +23,9 @@ type IntermediateReadWriteCloserWithAddr struct {
func (i *IntermediateReadWriteCloserWithAddr) Read(p []byte) (int, error) { func (i *IntermediateReadWriteCloserWithAddr) Read(p []byte) (int, error) {
return i.BufferedRead(p, func() error { return i.BufferedRead(p, func() error {
i.opts.QuickAck = false
i.opts.SimpleAck = false
buf := &bytes.Buffer{} buf := &bytes.Buffer{}
buf.Grow(4) buf.Grow(4)
+5 -6
View File
@@ -29,15 +29,16 @@ func (p *ProxyRequestReadWriteCloserWithAddr) Read(buf []byte) (int, error) {
return errors.Annotate(err, "Cannot read RPC tag") return errors.Annotate(err, "Cannot read RPC tag")
} }
if bytes.Equal(ansBuf.Bytes(), rpc.TagCloseExt) { switch {
case bytes.Equal(ansBuf.Bytes(), rpc.TagCloseExt):
return p.readCloseExt() return p.readCloseExt()
} else if bytes.Equal(ansBuf.Bytes(), rpc.TagProxyAns) { case bytes.Equal(ansBuf.Bytes(), rpc.TagProxyAns):
return p.readProxyAns(buf) return p.readProxyAns(buf)
} else if bytes.Equal(ansBuf.Bytes(), rpc.TagSimpleAck) { case bytes.Equal(ansBuf.Bytes(), rpc.TagSimpleAck):
return p.readSimpleAck() return p.readSimpleAck()
} }
return nil return errors.Errorf("Unknown RPC answer %s", ansBuf.Bytes())
}) })
} }
@@ -80,8 +81,6 @@ func (p *ProxyRequestReadWriteCloserWithAddr) Write(raw []byte) (int, error) {
if _, err := p.conn.Write(p.req.Bytes(raw)); err != nil { if _, err := p.conn.Write(p.req.Bytes(raw)); err != nil {
return 0, err return 0, err
} }
p.req.Options.SimpleAck = false
p.req.Options.QuickAck = false
return len(raw), nil return len(raw), nil
} }
+6 -5
View File
@@ -1,6 +1,7 @@
package wrappers package wrappers
import ( import (
"bytes"
"crypto/aes" "crypto/aes"
"crypto/cipher" "crypto/cipher"
"net" "net"
@@ -18,16 +19,16 @@ type BlockCipherReadWriteCloserWithAddr struct {
func (c *BlockCipherReadWriteCloserWithAddr) Read(p []byte) (int, error) { func (c *BlockCipherReadWriteCloserWithAddr) Read(p []byte) (int, error) {
return c.BufferedRead(p, func() error { return c.BufferedRead(p, func() error {
bufferLength := c.Buffer.Len() buf := &bytes.Buffer{}
for bufferLength%aes.BlockSize != 0 || bufferLength == 0 { for buf.Len()%aes.BlockSize != 0 || buf.Len() == 0 {
n, err := c.conn.Read(p) n, err := c.conn.Read(p)
if err != nil { if err != nil {
return errors.Annotate(err, "Cannot read from socket") return errors.Annotate(err, "Cannot read from socket")
} }
c.Buffer.Write(p[:n]) buf.Write(p[:n])
bufferLength = c.Buffer.Len()
} }
c.decryptor.CryptBlocks(c.Buffer.Bytes(), c.Buffer.Bytes()) c.decryptor.CryptBlocks(buf.Bytes(), buf.Bytes())
c.Buffer.Write(buf.Bytes())
return nil return nil
}) })