mirror of
https://github.com/xtaci/smux.git
synced 2024-04-21 10:51:48 +00:00
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
+22
-21
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user