diff --git a/frame.go b/frame.go index 1e92460..4dfaeab 100644 --- a/frame.go +++ b/frame.go @@ -61,6 +61,19 @@ func (f *Frame) UnmarshalBinary(bts []byte) error { return nil } +// ZeroCopyUnmarshal a byte slice into a frame, +// and just reference to the input slice +func (f *Frame) ZeroCopyUnmarshal(bts []byte) error { + f.ver = bts[0] + f.cmd = bts[1] + datalength := binary.LittleEndian.Uint16(bts[2:]) + f.sid = binary.LittleEndian.Uint32(bts[4:]) + if datalength > 0 { + f.data = bts[headerSize:] + } + return nil +} + type rawHeader []byte func (h rawHeader) Version() byte { diff --git a/frame_test.go b/frame_test.go index 81ded9c..ca8b687 100644 --- a/frame_test.go +++ b/frame_test.go @@ -17,10 +17,18 @@ func TestFrame(t *testing.T) { y.UnmarshalBinary(btsX) btsY, _ := y.MarshalBinary() + z := Frame{} + z.ZeroCopyUnmarshal(btsX) + btsZ, _ := z.MarshalBinary() + if !bytes.Equal(btsX, btsY) { t.Fatal("frame encode/decode failed") } + if !bytes.Equal(btsY, btsZ) { + t.Fatal("frame encode/decode zero copy failed") + } + t.Log(rawHeader(btsX).String()) t.Log(btsX) } diff --git a/session.go b/session.go index b0d2040..e3dfea5 100644 --- a/session.go +++ b/session.go @@ -1,6 +1,7 @@ package smux import ( + "bytes" "io" "sync" "sync/atomic" @@ -30,15 +31,14 @@ type Session struct { bucket int32 // token bucket bucketCond *sync.Cond // used for waiting for tokens - streamBuffers map[uint32][]Frame // stream input buffer - streams map[uint32]*Stream // all streams in this session - streamLock sync.Mutex // locks streams && frameQueues + streamBuffers map[uint32]*bytes.Buffer // stream input buffer + streams map[uint32]*Stream // all streams in this session + streamLock sync.Mutex // locks streams && frameQueues die chan struct{} // flag session has died dieLock sync.Mutex chAccepts chan *Stream chClosedStream chan uint32 - chBuferWriter chan []byte dataReady int32 // flag data has arrived } @@ -49,7 +49,7 @@ func newSession(config *Config, conn io.ReadWriteCloser, client bool) *Session { s.conn = conn s.config = config s.streams = make(map[uint32]*Stream) - s.streamBuffers = make(map[uint32][]Frame) + s.streamBuffers = make(map[uint32]*bytes.Buffer) s.chAccepts = make(chan *Stream, defaultAcceptBacklog) s.chClosedStream = make(chan uint32, defaultCloseWait) s.bucket = int32(config.MaxReceiveBuffer) @@ -73,9 +73,11 @@ func (s *Session) OpenStream() (*Stream, error) { sid := atomic.AddUint32(&s.nextStreamID, 2) stream := newStream(sid, s.config.MaxFrameSize, s) + streamBuffer := new(bytes.Buffer) s.streamLock.Lock() s.streams[sid] = stream + s.streamBuffers[sid] = streamBuffer s.streamLock.Unlock() s.writeFrame(newFrame(cmdSYN, sid)) @@ -145,19 +147,21 @@ func (s *Session) streamClosed(sid uint32) { } // nonblocking read from session pool, for streams -func (s *Session) nioread(sid uint32) (f Frame, ok bool) { +func (s *Session) nioread(sid uint32, p []byte) (n int) { s.streamLock.Lock() - if len(s.streamBuffers[sid]) > 0 { - f, ok = s.streamBuffers[sid][0], true - s.streamBuffers[sid] = s.streamBuffers[sid][1:] - atomic.AddInt32(&s.bucket, int32(len(f.data))) - s.bucketCond.Signal() + if streamBuffer, ok := s.streamBuffers[sid]; ok { + n, _ = streamBuffer.Read(p) + if n > 0 { + atomic.AddInt32(&s.bucket, int32(n)) + s.bucketCond.Signal() + } } s.streamLock.Unlock() return } // session read a frame from underlying connection +// it's data is pointed to the input buffer func (s *Session) readFrame(buffer []byte) (f Frame, err error) { if _, err := io.ReadFull(s.conn, buffer[:headerSize]); err != nil { return f, errors.Wrap(err, "readFrame") @@ -172,10 +176,10 @@ func (s *Session) readFrame(buffer []byte) (f Frame, err error) { if _, err := io.ReadFull(s.conn, buffer[headerSize:headerSize+length]); err != nil { return f, errors.Wrap(err, "readFrame") } - f.UnmarshalBinary(buffer[:headerSize+length]) + f.ZeroCopyUnmarshal(buffer[:headerSize+length]) return f, nil } - f.UnmarshalBinary(buffer[:headerSize]) + f.ZeroCopyUnmarshal(buffer[:headerSize]) return f, nil } @@ -186,13 +190,8 @@ func (s *Session) monitor() { case sid := <-s.chClosedStream: s.streamLock.Lock() delete(s.streams, sid) - streambuf := s.streamBuffers[sid] - var ntokens int - for k := range streambuf { - ntokens += len(streambuf[k].data) - } - if ntokens > 0 { - atomic.AddInt32(&s.bucket, int32(ntokens)) + if n := s.streamBuffers[sid].Len(); n > 0 { // return remaining tokens to the bucket + atomic.AddInt32(&s.bucket, int32(n)) s.bucketCond.Signal() } delete(s.streamBuffers, sid) @@ -226,7 +225,9 @@ func (s *Session) recvLoop() { s.streamLock.Lock() if _, ok := s.streams[f.sid]; !ok { stream := newStream(f.sid, s.config.MaxFrameSize, s) + streamBuffer := new(bytes.Buffer) s.streams[f.sid] = stream + s.streamBuffers[f.sid] = streamBuffer go func() { s.chAccepts <- stream }() } else { // stream exists, RST the peer go s.writeFrame(newFrame(cmdRST, f.sid)) @@ -243,7 +244,7 @@ func (s *Session) recvLoop() { s.streamLock.Lock() if stream, ok := s.streams[f.sid]; ok { atomic.AddInt32(&s.bucket, -int32(len(f.data))) - s.streamBuffers[f.sid] = append(s.streamBuffers[f.sid], f) + s.streamBuffers[f.sid].Write(f.data) stream.notifyReadEvent() } else { // stream is absent go s.writeFrame(newFrame(cmdRST, f.sid)) diff --git a/stream.go b/stream.go index f6d9f58..17d6717 100644 --- a/stream.go +++ b/stream.go @@ -11,8 +11,6 @@ import ( // Stream implements io.ReadWriteCloser type Stream struct { id uint32 - buffer []byte - rlock sync.Mutex rstflag int32 sess *Session frameSize int @@ -32,10 +30,8 @@ func newStream(id uint32, frameSize int, sess *Session) *Stream { return s } -// Read implements io.ReadWriteCloser, only one reader can enter +// Read implements io.ReadWriteCloser func (s *Stream) Read(b []byte) (n int, err error) { - s.rlock.Lock() - defer s.rlock.Unlock() READ: select { case <-s.die: @@ -43,16 +39,8 @@ READ: default: } - if len(s.buffer) > 0 { - n, err = copy(b, s.buffer), nil - s.buffer = s.buffer[n:] - return - } - - if f, ok := s.sess.nioread(s.id); ok { - n, err = copy(b, f.data), nil - s.buffer = f.data[n:] - return + if n = s.sess.nioread(s.id, b); n > 0 { + return n, err } else if atomic.LoadInt32(&s.rstflag) == 1 { s.Close() return 0, errors.New(errBrokenPipe)