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:
xtaci
2019-12-26 12:06:21 +08:00
parent 6125539908
commit f386d90508
7 changed files with 174 additions and 61 deletions
+1 -2
View File
@@ -1,9 +1,8 @@
package smux package smux
import ( import (
"errors"
"sync" "sync"
"github.com/pkg/errors"
) )
var defaultAllocator *Allocator var defaultAllocator *Allocator
-2
View File
@@ -1,5 +1,3 @@
module github.com/xtaci/smux module github.com/xtaci/smux
go 1.13 go 1.13
require github.com/pkg/errors v0.8.1
-2
View File
@@ -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=
+1 -2
View File
@@ -5,11 +5,10 @@
package smux package smux
import ( import (
"errors"
"fmt" "fmt"
"io" "io"
"time" "time"
"github.com/pkg/errors"
) )
// Config is used to tune the Smux session // Config is used to tune the Smux session
+77 -23
View File
@@ -9,7 +9,7 @@ import (
"sync/atomic" "sync/atomic"
"time" "time"
"github.com/pkg/errors" "errors"
) )
const ( const (
@@ -17,9 +17,11 @@ const (
) )
var ( var (
ErrInvalidProtocol = errors.New("invalid protocol") ErrInvalidProtocol = errors.New("invalid protocol")
ErrGoAway = errors.New("stream id overflows, should start a new connection") ErrGoAway = errors.New("stream id overflows, should start a new connection")
ErrTimeout = errors.New("timeout") ErrTimeout = errors.New("timeout")
ErrInvalidOperation = errors.New("invalid parameters on poll")
ErrWouldBlock = errors.New("operation would block on IO")
) )
type writeRequest struct { type writeRequest struct {
@@ -77,6 +79,12 @@ type Session struct {
shaper chan writeRequest // a shaper for writing shaper chan writeRequest // a shaper for writing
writes chan writeRequest 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 { 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.chSocketReadError = make(chan struct{})
s.chSocketWriteError = make(chan struct{}) s.chSocketWriteError = make(chan struct{})
s.chProtoError = make(chan struct{}) s.chProtoError = make(chan struct{})
s.chPollEventNotify = make(chan struct{}, 1)
s.pollEvents = make(map[uint32]*Stream)
if client { if client {
s.nextStreamID = 1 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 // OpenStream is used to create a new stream
func (s *Session) OpenStream() (*Stream, error) { func (s *Session) OpenStream() (*Stream, error) {
if s.IsClosed() { if s.IsClosed() {
return nil, errors.WithStack(io.ErrClosedPipe) return nil, io.ErrClosedPipe
} }
// generate stream id // generate stream id
s.nextStreamIDLock.Lock() s.nextStreamIDLock.Lock()
if s.goAway > 0 { if s.goAway > 0 {
s.nextStreamIDLock.Unlock() s.nextStreamIDLock.Unlock()
return nil, errors.WithStack(ErrGoAway) return nil, ErrGoAway
} }
s.nextStreamID += 2 s.nextStreamID += 2
@@ -125,14 +135,14 @@ func (s *Session) OpenStream() (*Stream, error) {
if sid == sid%2 { // stream-id overflows if sid == sid%2 { // stream-id overflows
s.goAway = 1 s.goAway = 1
s.nextStreamIDLock.Unlock() s.nextStreamIDLock.Unlock()
return nil, errors.WithStack(ErrGoAway) return nil, ErrGoAway
} }
s.nextStreamIDLock.Unlock() s.nextStreamIDLock.Unlock()
stream := newStream(sid, s.config.MaxFrameSize, s) stream := newStream(sid, s.config.MaxFrameSize, s)
if _, err := s.writeFrame(newFrame(cmdSYN, sid)); err != nil { if _, err := s.writeFrame(newFrame(cmdSYN, sid)); err != nil {
return nil, errors.WithStack(err) return nil, err
} }
s.streamLock.Lock() s.streamLock.Lock()
@@ -141,7 +151,7 @@ func (s *Session) OpenStream() (*Stream, error) {
case <-s.chSocketWriteError: case <-s.chSocketWriteError:
return nil, s.socketWriteError.Load().(error) return nil, s.socketWriteError.Load().(error)
case <-s.die: case <-s.die:
return nil, errors.WithStack(io.ErrClosedPipe) return nil, io.ErrClosedPipe
default: default:
s.streams[sid] = stream s.streams[sid] = stream
return stream, nil return stream, nil
@@ -167,13 +177,13 @@ func (s *Session) AcceptStream() (*Stream, error) {
case stream := <-s.chAccepts: case stream := <-s.chAccepts:
return stream, nil return stream, nil
case <-deadline: case <-deadline:
return nil, errors.WithStack(ErrTimeout) return nil, ErrTimeout
case <-s.chSocketReadError: case <-s.chSocketReadError:
return nil, s.socketReadError.Load().(error) return nil, s.socketReadError.Load().(error)
case <-s.chProtoError: case <-s.chProtoError:
return nil, s.protoError.Load().(error) return nil, s.protoError.Load().(error)
case <-s.die: 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() s.streamLock.Unlock()
return s.conn.Close() return s.conn.Close()
} else { } 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) delete(s.streams, sid)
s.streamLock.Unlock() 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 // 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 { if _, err := io.ReadFull(s.conn, hdr[:]); err == nil {
atomic.StoreInt32(&s.dataReady, 1) atomic.StoreInt32(&s.dataReady, 1)
if hdr.Version() != version { if hdr.Version() != version {
s.notifyProtoError(errors.WithStack(ErrInvalidProtocol)) s.notifyProtoError(ErrInvalidProtocol)
return return
} }
sid := hdr.StreamID() sid := hdr.StreamID()
@@ -350,16 +404,16 @@ func (s *Session) recvLoop() {
} }
s.streamLock.Unlock() s.streamLock.Unlock()
} else { } else {
s.notifyReadError(errors.WithStack(err)) s.notifyReadError(err)
return return
} }
} }
default: default:
s.notifyProtoError(errors.WithStack(ErrInvalidProtocol)) s.notifyProtoError(ErrInvalidProtocol)
return return
} }
} else { } else {
s.notifyReadError(errors.WithStack(err)) s.notifyReadError(err)
return return
} }
} }
@@ -453,7 +507,7 @@ func (s *Session) sendLoop() {
result := writeResult{ result := writeResult{
n: n, n: n,
err: errors.WithStack(err), err: err,
} }
request.result <- result request.result <- result
@@ -461,7 +515,7 @@ func (s *Session) sendLoop() {
// store conn error // store conn error
if err != nil { if err != nil {
s.notifyWriteError(errors.WithStack(err)) s.notifyWriteError(err)
return return
} }
} }
@@ -484,21 +538,21 @@ func (s *Session) writeFrameInternal(f Frame, deadline <-chan time.Time, prio ui
select { select {
case s.shaper <- req: case s.shaper <- req:
case <-s.die: case <-s.die:
return 0, errors.WithStack(io.ErrClosedPipe) return 0, io.ErrClosedPipe
case <-s.chSocketWriteError: case <-s.chSocketWriteError:
return 0, s.socketWriteError.Load().(error) return 0, s.socketWriteError.Load().(error)
case <-deadline: case <-deadline:
return 0, errors.WithStack(ErrTimeout) return 0, ErrTimeout
} }
select { select {
case result := <-req.result: case result := <-req.result:
return result.n, errors.WithStack(result.err) return result.n, result.err
case <-s.die: case <-s.die:
return 0, errors.WithStack(io.ErrClosedPipe) return 0, io.ErrClosedPipe
case <-s.chSocketWriteError: case <-s.chSocketWriteError:
return 0, s.socketWriteError.Load().(error) return 0, s.socketWriteError.Load().(error)
case <-deadline: case <-deadline:
return 0, errors.WithStack(ErrTimeout) return 0, ErrTimeout
} }
} }
+48
View File
@@ -95,6 +95,54 @@ func TestEcho(t *testing.T) {
session.Close() 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) { func TestWriteTo(t *testing.T) {
const N = 1 << 20 const N = 1 << 20
// server // server
+47 -30
View File
@@ -6,8 +6,6 @@ import (
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
"github.com/pkg/errors"
) )
// Stream implements net.Conn // Stream implements net.Conn
@@ -59,33 +57,48 @@ func (s *Stream) ID() uint32 {
// Read implements net.Conn // Read implements net.Conn
func (s *Stream) Read(b []byte) (n int, err error) { 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 { if len(b) == 0 {
return 0, nil return 0, nil
} }
for { s.bufferLock.Lock()
s.bufferLock.Lock() if len(s.buffers) > 0 {
if len(s.buffers) > 0 { n = copy(b, s.buffers[0])
n = copy(b, s.buffers[0]) s.buffers[0] = s.buffers[0][n:]
s.buffers[0] = s.buffers[0][n:] if len(s.buffers[0]) == 0 {
if len(s.buffers[0]) == 0 { s.buffers[0] = nil
s.buffers[0] = nil s.buffers = s.buffers[1:]
s.buffers = s.buffers[1:] // full recycle
// full recycle defaultAllocator.Put(s.heads[0])
defaultAllocator.Put(s.heads[0]) s.heads = s.heads[1:]
s.heads = s.heads[1:]
}
} }
s.bufferLock.Unlock() }
s.bufferLock.Unlock()
if n > 0 { if n > 0 {
s.sess.returnTokens(n) s.sess.returnTokens(n)
return n, nil return n, nil
} }
if ew := s.waitRead(); ew != nil { select {
return 0, ew case <-s.die:
} return 0, io.EOF
default:
return 0, ErrWouldBlock
} }
} }
@@ -131,15 +144,15 @@ func (s *Stream) waitRead() error {
case <-s.chReadEvent: case <-s.chReadEvent:
return nil return nil
case <-s.chFinEvent: case <-s.chFinEvent:
return errors.WithStack(io.EOF) return io.EOF
case <-s.sess.chSocketReadError: case <-s.sess.chSocketReadError:
return s.sess.socketReadError.Load().(error) return s.sess.socketReadError.Load().(error)
case <-s.sess.chProtoError: case <-s.sess.chProtoError:
return s.sess.protoError.Load().(error) return s.sess.protoError.Load().(error)
case <-deadline: case <-deadline:
return errors.WithStack(ErrTimeout) return ErrTimeout
case <-s.die: 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 // check if stream has closed
select { select {
case <-s.die: case <-s.die:
return 0, errors.WithStack(io.ErrClosedPipe) return 0, io.ErrClosedPipe
default: default:
} }
@@ -175,7 +188,7 @@ func (s *Stream) Write(b []byte) (n int, err error) {
s.numWrite++ s.numWrite++
sent += n sent += n
if err != nil { 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) s.sess.streamClosed(s.id)
return err return err
} else { } 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. // A zero time value disables the deadlines.
func (s *Stream) SetDeadline(t time.Time) error { func (s *Stream) SetDeadline(t time.Time) error {
if err := s.SetReadDeadline(t); err != nil { if err := s.SetReadDeadline(t); err != nil {
return errors.WithStack(err) return err
} }
if err := s.SetWriteDeadline(t); err != nil { if err := s.SetWriteDeadline(t); err != nil {
return errors.WithStack(err) return err
} }
return nil return nil
} }
@@ -263,6 +276,10 @@ func (s *Stream) pushBytes(buf []byte) (written int, err error) {
s.bufferLock.Lock() s.bufferLock.Lock()
s.buffers = append(s.buffers, buf) s.buffers = append(s.buffers, buf)
s.heads = append(s.heads, buf) s.heads = append(s.heads, buf)
// Edge-Trigger
if len(s.buffers) == 1 {
s.sess.notifyPoll(s)
}
s.bufferLock.Unlock() s.bufferLock.Unlock()
return return
} }