diff --git a/frame.go b/frame.go index fad0e34..338a259 100644 --- a/frame.go +++ b/frame.go @@ -63,7 +63,7 @@ func (f *Frame) UnmarshalBinary(bts []byte) error { // zeroCopyUnmarshal a byte slice into a frame, // and just reference to the input slice -func (f *Frame) zeroCopyUnmarshal(bts []byte) error { +func (f *Frame) zeroCopyUnmarshal(bts []byte) { f.ver = bts[0] f.cmd = bts[1] datalength := binary.LittleEndian.Uint16(bts[2:]) @@ -71,7 +71,6 @@ func (f *Frame) zeroCopyUnmarshal(bts []byte) error { if datalength > 0 { f.data = bts[headerSize:] } - return nil } type rawHeader []byte diff --git a/session.go b/session.go index d31ff17..76cf5b3 100644 --- a/session.go +++ b/session.go @@ -15,6 +15,7 @@ const ( const ( errBrokenPipe = "broken pipe" + errConnReset = "connection reset by peer" errInvalidProtocol = "invalid protocol version" ) @@ -67,11 +68,13 @@ func (s *Session) OpenStream() (*Stream, error) { sid := atomic.AddUint32(&s.nextStreamID, 2) stream := newStream(sid, s.config.MaxFrameSize, s) + if _, err := s.writeFrame(newFrame(cmdSYN, sid)); err != nil { + return nil, errors.Wrap(err, "writeFrame") + } + s.streamLock.Lock() s.streams[sid] = stream s.streamLock.Unlock() - - s.writeFrame(newFrame(cmdSYN, sid)) return stream, nil } @@ -87,7 +90,7 @@ func (s *Session) AcceptStream() (*Stream, error) { } // Close is used to close the session and all streams. -func (s *Session) Close() error { +func (s *Session) Close() (err error) { s.dieLock.Lock() defer s.dieLock.Unlock() @@ -95,16 +98,15 @@ func (s *Session) Close() error { case <-s.die: return errors.New(errBrokenPipe) default: + close(s.die) s.streamLock.Lock() for k := range s.streams { s.streams[k].sessionClose() } s.streamLock.Unlock() - s.conn.Close() - close(s.die) s.bucketCond.Signal() + return s.conn.Close() } - return nil } // IsClosed does a safe check to see if we have shutdown diff --git a/stream.go b/stream.go index a66bb37..6bb9b00 100644 --- a/stream.go +++ b/stream.go @@ -49,8 +49,8 @@ READ: s.sess.returnTokens(n) return n, nil } else if atomic.LoadInt32(&s.rstflag) == 1 { - s.Close() - return 0, errors.New(errBrokenPipe) + _ = s.Close() + return 0, errors.New(errConnReset) } select { @@ -94,9 +94,9 @@ func (s *Stream) Close() error { default: close(s.die) s.sess.streamClosed(s.id) - s.sess.writeFrame(newFrame(cmdRST, s.id)) + _, err := s.sess.writeFrame(newFrame(cmdRST, s.id)) + return err } - return nil } // session closes the stream