mirror of
https://github.com/xtaci/smux.git
synced 2024-04-21 10:51:48 +00:00
Squashed commit of the following:
commit 486431f53deb10b0e6ee51b47785f6f06f0d5733 Author: xtaci <daniel820313@gmail.com> Date: Thu Dec 26 12:02:52 2019 +0800 remove in-package pkg/errors commit 213597603a726735678024ec6051800c46ed0bd0 Author: xtaci <daniel820313@gmail.com> Date: Thu Dec 26 11:58:36 2019 +0800 fix poll commit d1f7aeefbeb40da99d47ba2b7db5bd403c09ddba Author: xtaci <daniel820313@gmail.com> Date: Thu Dec 26 11:17:28 2019 +0800 fix error commit 152918f9b531c06b42c764f0aff35d4bb9315713 Author: xtaci <daniel820313@gmail.com> Date: Wed Dec 25 23:27:18 2019 +0800 poll for v1
This commit is contained in:
@@ -1,9 +1,8 @@
|
||||
package smux
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sync"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
var defaultAllocator *Allocator
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
module github.com/xtaci/smux
|
||||
|
||||
go 1.13
|
||||
|
||||
require github.com/pkg/errors v0.8.1
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
github.com/pkg/errors v0.8.1 h1:iURUrRGxPUNPdy5/HRSm+Yj6okJ6UtLINN0Q9M4+h3I=
|
||||
github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||
|
||||
@@ -5,11 +5,10 @@
|
||||
package smux
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// Config is used to tune the Smux session
|
||||
|
||||
+77
-23
@@ -9,7 +9,7 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"errors"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -17,9 +17,11 @@ const (
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidProtocol = errors.New("invalid protocol")
|
||||
ErrGoAway = errors.New("stream id overflows, should start a new connection")
|
||||
ErrTimeout = errors.New("timeout")
|
||||
ErrInvalidProtocol = errors.New("invalid protocol")
|
||||
ErrGoAway = errors.New("stream id overflows, should start a new connection")
|
||||
ErrTimeout = errors.New("timeout")
|
||||
ErrInvalidOperation = errors.New("invalid parameters on poll")
|
||||
ErrWouldBlock = errors.New("operation would block on IO")
|
||||
)
|
||||
|
||||
type writeRequest struct {
|
||||
@@ -77,6 +79,12 @@ type Session struct {
|
||||
|
||||
shaper chan writeRequest // a shaper for writing
|
||||
writes chan writeRequest
|
||||
|
||||
// Edge-Triggered PollIn support
|
||||
// Streams which become 'readable', will return from PollWait()
|
||||
pollEvents map[uint32]*Stream
|
||||
pollEventsLock sync.Mutex
|
||||
chPollEventNotify chan struct{} // notify new events
|
||||
}
|
||||
|
||||
func newSession(config *Config, conn io.ReadWriteCloser, client bool) *Session {
|
||||
@@ -93,6 +101,8 @@ func newSession(config *Config, conn io.ReadWriteCloser, client bool) *Session {
|
||||
s.chSocketReadError = make(chan struct{})
|
||||
s.chSocketWriteError = make(chan struct{})
|
||||
s.chProtoError = make(chan struct{})
|
||||
s.chPollEventNotify = make(chan struct{}, 1)
|
||||
s.pollEvents = make(map[uint32]*Stream)
|
||||
|
||||
if client {
|
||||
s.nextStreamID = 1
|
||||
@@ -110,14 +120,14 @@ func newSession(config *Config, conn io.ReadWriteCloser, client bool) *Session {
|
||||
// OpenStream is used to create a new stream
|
||||
func (s *Session) OpenStream() (*Stream, error) {
|
||||
if s.IsClosed() {
|
||||
return nil, errors.WithStack(io.ErrClosedPipe)
|
||||
return nil, io.ErrClosedPipe
|
||||
}
|
||||
|
||||
// generate stream id
|
||||
s.nextStreamIDLock.Lock()
|
||||
if s.goAway > 0 {
|
||||
s.nextStreamIDLock.Unlock()
|
||||
return nil, errors.WithStack(ErrGoAway)
|
||||
return nil, ErrGoAway
|
||||
}
|
||||
|
||||
s.nextStreamID += 2
|
||||
@@ -125,14 +135,14 @@ func (s *Session) OpenStream() (*Stream, error) {
|
||||
if sid == sid%2 { // stream-id overflows
|
||||
s.goAway = 1
|
||||
s.nextStreamIDLock.Unlock()
|
||||
return nil, errors.WithStack(ErrGoAway)
|
||||
return nil, ErrGoAway
|
||||
}
|
||||
s.nextStreamIDLock.Unlock()
|
||||
|
||||
stream := newStream(sid, s.config.MaxFrameSize, s)
|
||||
|
||||
if _, err := s.writeFrame(newFrame(cmdSYN, sid)); err != nil {
|
||||
return nil, errors.WithStack(err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.streamLock.Lock()
|
||||
@@ -141,7 +151,7 @@ func (s *Session) OpenStream() (*Stream, error) {
|
||||
case <-s.chSocketWriteError:
|
||||
return nil, s.socketWriteError.Load().(error)
|
||||
case <-s.die:
|
||||
return nil, errors.WithStack(io.ErrClosedPipe)
|
||||
return nil, io.ErrClosedPipe
|
||||
default:
|
||||
s.streams[sid] = stream
|
||||
return stream, nil
|
||||
@@ -167,13 +177,13 @@ func (s *Session) AcceptStream() (*Stream, error) {
|
||||
case stream := <-s.chAccepts:
|
||||
return stream, nil
|
||||
case <-deadline:
|
||||
return nil, errors.WithStack(ErrTimeout)
|
||||
return nil, ErrTimeout
|
||||
case <-s.chSocketReadError:
|
||||
return nil, s.socketReadError.Load().(error)
|
||||
case <-s.chProtoError:
|
||||
return nil, s.protoError.Load().(error)
|
||||
case <-s.die:
|
||||
return nil, errors.WithStack(io.ErrClosedPipe)
|
||||
return nil, io.ErrClosedPipe
|
||||
}
|
||||
}
|
||||
|
||||
@@ -198,7 +208,46 @@ func (s *Session) Close() error {
|
||||
s.streamLock.Unlock()
|
||||
return s.conn.Close()
|
||||
} else {
|
||||
return errors.WithStack(io.ErrClosedPipe)
|
||||
return io.ErrClosedPipe
|
||||
}
|
||||
}
|
||||
|
||||
// PollWait returns streams which became readable
|
||||
func (s *Session) PollWait(events []*Stream) (int, error) {
|
||||
if len(events) == 0 {
|
||||
return -1, ErrInvalidOperation
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-s.chPollEventNotify:
|
||||
s.pollEventsLock.Lock()
|
||||
i := 0
|
||||
for id, stream := range s.pollEvents {
|
||||
if i >= len(events) {
|
||||
break
|
||||
}
|
||||
events[i] = stream
|
||||
i++
|
||||
delete(s.pollEvents, id)
|
||||
}
|
||||
s.pollEventsLock.Unlock()
|
||||
return i, nil
|
||||
case <-s.die:
|
||||
return -1, io.ErrClosedPipe
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// streams notify session pollin events
|
||||
func (s *Session) notifyPoll(stream *Stream) {
|
||||
s.pollEventsLock.Lock()
|
||||
s.pollEvents[stream.id] = stream
|
||||
s.pollEventsLock.Unlock()
|
||||
|
||||
select {
|
||||
case s.chPollEventNotify <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
@@ -288,6 +337,11 @@ func (s *Session) streamClosed(sid uint32) {
|
||||
}
|
||||
delete(s.streams, sid)
|
||||
s.streamLock.Unlock()
|
||||
|
||||
// poll remove
|
||||
s.pollEventsLock.Lock()
|
||||
delete(s.pollEvents, sid)
|
||||
s.pollEventsLock.Unlock()
|
||||
}
|
||||
|
||||
// returnTokens is called by stream to return token after read
|
||||
@@ -314,7 +368,7 @@ func (s *Session) recvLoop() {
|
||||
if _, err := io.ReadFull(s.conn, hdr[:]); err == nil {
|
||||
atomic.StoreInt32(&s.dataReady, 1)
|
||||
if hdr.Version() != version {
|
||||
s.notifyProtoError(errors.WithStack(ErrInvalidProtocol))
|
||||
s.notifyProtoError(ErrInvalidProtocol)
|
||||
return
|
||||
}
|
||||
sid := hdr.StreamID()
|
||||
@@ -350,16 +404,16 @@ func (s *Session) recvLoop() {
|
||||
}
|
||||
s.streamLock.Unlock()
|
||||
} else {
|
||||
s.notifyReadError(errors.WithStack(err))
|
||||
s.notifyReadError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
default:
|
||||
s.notifyProtoError(errors.WithStack(ErrInvalidProtocol))
|
||||
s.notifyProtoError(ErrInvalidProtocol)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
s.notifyReadError(errors.WithStack(err))
|
||||
s.notifyReadError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -453,7 +507,7 @@ func (s *Session) sendLoop() {
|
||||
|
||||
result := writeResult{
|
||||
n: n,
|
||||
err: errors.WithStack(err),
|
||||
err: err,
|
||||
}
|
||||
|
||||
request.result <- result
|
||||
@@ -461,7 +515,7 @@ func (s *Session) sendLoop() {
|
||||
|
||||
// store conn error
|
||||
if err != nil {
|
||||
s.notifyWriteError(errors.WithStack(err))
|
||||
s.notifyWriteError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -484,21 +538,21 @@ func (s *Session) writeFrameInternal(f Frame, deadline <-chan time.Time, prio ui
|
||||
select {
|
||||
case s.shaper <- req:
|
||||
case <-s.die:
|
||||
return 0, errors.WithStack(io.ErrClosedPipe)
|
||||
return 0, io.ErrClosedPipe
|
||||
case <-s.chSocketWriteError:
|
||||
return 0, s.socketWriteError.Load().(error)
|
||||
case <-deadline:
|
||||
return 0, errors.WithStack(ErrTimeout)
|
||||
return 0, ErrTimeout
|
||||
}
|
||||
|
||||
select {
|
||||
case result := <-req.result:
|
||||
return result.n, errors.WithStack(result.err)
|
||||
return result.n, result.err
|
||||
case <-s.die:
|
||||
return 0, errors.WithStack(io.ErrClosedPipe)
|
||||
return 0, io.ErrClosedPipe
|
||||
case <-s.chSocketWriteError:
|
||||
return 0, s.socketWriteError.Load().(error)
|
||||
case <-deadline:
|
||||
return 0, errors.WithStack(ErrTimeout)
|
||||
return 0, ErrTimeout
|
||||
}
|
||||
}
|
||||
|
||||
@@ -95,6 +95,54 @@ func TestEcho(t *testing.T) {
|
||||
session.Close()
|
||||
}
|
||||
|
||||
func TestPoll(t *testing.T) {
|
||||
_, stop, cli, err := setupServer(t)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer stop()
|
||||
session, _ := Client(cli, nil)
|
||||
stream, _ := session.OpenStream()
|
||||
|
||||
const N = 100
|
||||
var received int
|
||||
|
||||
tx := make([]byte, 128)
|
||||
go func() {
|
||||
for i := 0; i < N; i++ {
|
||||
stream.Write(tx)
|
||||
}
|
||||
}()
|
||||
|
||||
buf := make([]byte, 128)
|
||||
events := make([]*Stream, 128)
|
||||
for {
|
||||
n, err := session.PollWait(events)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
stream := events[i]
|
||||
for {
|
||||
n, err := stream.TryRead(buf)
|
||||
if err == ErrWouldBlock {
|
||||
break
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
received += n
|
||||
if received == len(tx)*N {
|
||||
session.Close()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteTo(t *testing.T) {
|
||||
const N = 1 << 20
|
||||
// server
|
||||
|
||||
@@ -6,8 +6,6 @@ import (
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// Stream implements net.Conn
|
||||
@@ -59,33 +57,48 @@ func (s *Stream) ID() uint32 {
|
||||
|
||||
// Read implements net.Conn
|
||||
func (s *Stream) Read(b []byte) (n int, err error) {
|
||||
for {
|
||||
n, err = s.TryRead(b)
|
||||
if err == ErrWouldBlock {
|
||||
if ew := s.waitRead(); ew != nil {
|
||||
return 0, ew
|
||||
}
|
||||
} else {
|
||||
return n, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TryRead is the nonblocking version of Read
|
||||
func (s *Stream) TryRead(b []byte) (n int, err error) {
|
||||
if len(b) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
for {
|
||||
s.bufferLock.Lock()
|
||||
if len(s.buffers) > 0 {
|
||||
n = copy(b, s.buffers[0])
|
||||
s.buffers[0] = s.buffers[0][n:]
|
||||
if len(s.buffers[0]) == 0 {
|
||||
s.buffers[0] = nil
|
||||
s.buffers = s.buffers[1:]
|
||||
// full recycle
|
||||
defaultAllocator.Put(s.heads[0])
|
||||
s.heads = s.heads[1:]
|
||||
}
|
||||
s.bufferLock.Lock()
|
||||
if len(s.buffers) > 0 {
|
||||
n = copy(b, s.buffers[0])
|
||||
s.buffers[0] = s.buffers[0][n:]
|
||||
if len(s.buffers[0]) == 0 {
|
||||
s.buffers[0] = nil
|
||||
s.buffers = s.buffers[1:]
|
||||
// full recycle
|
||||
defaultAllocator.Put(s.heads[0])
|
||||
s.heads = s.heads[1:]
|
||||
}
|
||||
s.bufferLock.Unlock()
|
||||
}
|
||||
s.bufferLock.Unlock()
|
||||
|
||||
if n > 0 {
|
||||
s.sess.returnTokens(n)
|
||||
return n, nil
|
||||
}
|
||||
if n > 0 {
|
||||
s.sess.returnTokens(n)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
if ew := s.waitRead(); ew != nil {
|
||||
return 0, ew
|
||||
}
|
||||
select {
|
||||
case <-s.die:
|
||||
return 0, io.EOF
|
||||
default:
|
||||
return 0, ErrWouldBlock
|
||||
}
|
||||
}
|
||||
|
||||
@@ -131,15 +144,15 @@ func (s *Stream) waitRead() error {
|
||||
case <-s.chReadEvent:
|
||||
return nil
|
||||
case <-s.chFinEvent:
|
||||
return errors.WithStack(io.EOF)
|
||||
return io.EOF
|
||||
case <-s.sess.chSocketReadError:
|
||||
return s.sess.socketReadError.Load().(error)
|
||||
case <-s.sess.chProtoError:
|
||||
return s.sess.protoError.Load().(error)
|
||||
case <-deadline:
|
||||
return errors.WithStack(ErrTimeout)
|
||||
return ErrTimeout
|
||||
case <-s.die:
|
||||
return errors.WithStack(io.ErrClosedPipe)
|
||||
return io.ErrClosedPipe
|
||||
}
|
||||
|
||||
}
|
||||
@@ -156,7 +169,7 @@ func (s *Stream) Write(b []byte) (n int, err error) {
|
||||
// check if stream has closed
|
||||
select {
|
||||
case <-s.die:
|
||||
return 0, errors.WithStack(io.ErrClosedPipe)
|
||||
return 0, io.ErrClosedPipe
|
||||
default:
|
||||
}
|
||||
|
||||
@@ -175,7 +188,7 @@ func (s *Stream) Write(b []byte) (n int, err error) {
|
||||
s.numWrite++
|
||||
sent += n
|
||||
if err != nil {
|
||||
return sent, errors.WithStack(err)
|
||||
return sent, err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -196,7 +209,7 @@ func (s *Stream) Close() error {
|
||||
s.sess.streamClosed(s.id)
|
||||
return err
|
||||
} else {
|
||||
return errors.WithStack(io.ErrClosedPipe)
|
||||
return io.ErrClosedPipe
|
||||
}
|
||||
}
|
||||
|
||||
@@ -227,10 +240,10 @@ func (s *Stream) SetWriteDeadline(t time.Time) error {
|
||||
// A zero time value disables the deadlines.
|
||||
func (s *Stream) SetDeadline(t time.Time) error {
|
||||
if err := s.SetReadDeadline(t); err != nil {
|
||||
return errors.WithStack(err)
|
||||
return err
|
||||
}
|
||||
if err := s.SetWriteDeadline(t); err != nil {
|
||||
return errors.WithStack(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -263,6 +276,10 @@ func (s *Stream) pushBytes(buf []byte) (written int, err error) {
|
||||
s.bufferLock.Lock()
|
||||
s.buffers = append(s.buffers, buf)
|
||||
s.heads = append(s.heads, buf)
|
||||
// Edge-Trigger
|
||||
if len(s.buffers) == 1 {
|
||||
s.sess.notifyPoll(s)
|
||||
}
|
||||
s.bufferLock.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user