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
import (
"errors"
"sync"
"github.com/pkg/errors"
)
var defaultAllocator *Allocator
-2
View File
@@ -1,5 +1,3 @@
module github.com/xtaci/smux
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
import (
"errors"
"fmt"
"io"
"time"
"github.com/pkg/errors"
)
// Config is used to tune the Smux session
+77 -23
View File
@@ -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
}
}
+48
View File
@@ -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
+47 -30
View File
@@ -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
}