mirror of
https://github.com/xtaci/smux.git
synced 2024-04-21 10:51:48 +00:00
+79
-11
@@ -19,6 +19,16 @@ const (
|
||||
errInvalidProtocol = "invalid protocol version"
|
||||
)
|
||||
|
||||
type writeRequest struct {
|
||||
frame Frame
|
||||
result chan writeResult
|
||||
}
|
||||
|
||||
type writeResult struct {
|
||||
n int
|
||||
err error
|
||||
}
|
||||
|
||||
// Session defines a multiplexed connection for streams
|
||||
type Session struct {
|
||||
conn io.ReadWriteCloser
|
||||
@@ -39,6 +49,10 @@ type Session struct {
|
||||
|
||||
xmitPool sync.Pool
|
||||
dataReady int32 // flag data has arrived
|
||||
|
||||
deadline atomic.Value
|
||||
|
||||
writes chan writeRequest
|
||||
}
|
||||
|
||||
func newSession(config *Config, conn io.ReadWriteCloser, client bool) *Session {
|
||||
@@ -53,6 +67,7 @@ func newSession(config *Config, conn io.ReadWriteCloser, client bool) *Session {
|
||||
s.xmitPool.New = func() interface{} {
|
||||
return make([]byte, (1<<16)+headerSize)
|
||||
}
|
||||
s.writes = make(chan writeRequest)
|
||||
|
||||
if client {
|
||||
s.nextStreamID = 1
|
||||
@@ -60,6 +75,7 @@ func newSession(config *Config, conn io.ReadWriteCloser, client bool) *Session {
|
||||
s.nextStreamID = 2
|
||||
}
|
||||
go s.recvLoop()
|
||||
go s.sendLoop()
|
||||
go s.keepalive()
|
||||
return s
|
||||
}
|
||||
@@ -86,9 +102,17 @@ func (s *Session) OpenStream() (*Stream, error) {
|
||||
// AcceptStream is used to block until the next available stream
|
||||
// is ready to be accepted.
|
||||
func (s *Session) AcceptStream() (*Stream, error) {
|
||||
var deadline <-chan time.Time
|
||||
if d, ok := s.deadline.Load().(time.Time); ok && !d.IsZero() {
|
||||
timer := time.NewTimer(d.Sub(time.Now()))
|
||||
defer timer.Stop()
|
||||
deadline = timer.C
|
||||
}
|
||||
select {
|
||||
case stream := <-s.chAccepts:
|
||||
return stream, nil
|
||||
case <-deadline:
|
||||
return nil, errTimeout
|
||||
case <-s.die:
|
||||
return nil, errors.New(errBrokenPipe)
|
||||
}
|
||||
@@ -134,6 +158,13 @@ func (s *Session) NumStreams() int {
|
||||
return len(s.streams)
|
||||
}
|
||||
|
||||
// SetDeadline sets a deadline used by Accept* calls.
|
||||
// A zero time value disables the deadline.
|
||||
func (s *Session) SetDeadline(t time.Time) error {
|
||||
s.deadline.Store(t)
|
||||
return nil
|
||||
}
|
||||
|
||||
// notify the session that a stream has closed
|
||||
func (s *Session) streamClosed(sid uint32) {
|
||||
s.streamLock.Lock()
|
||||
@@ -257,19 +288,56 @@ func (s *Session) keepalive() {
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Session) sendLoop() {
|
||||
for {
|
||||
select {
|
||||
case <-s.die:
|
||||
return
|
||||
case request, ok := <-s.writes:
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
buf := s.xmitPool.Get().([]byte)
|
||||
buf[0] = request.frame.ver
|
||||
buf[1] = request.frame.cmd
|
||||
binary.LittleEndian.PutUint16(buf[2:], uint16(len(request.frame.data)))
|
||||
binary.LittleEndian.PutUint32(buf[4:], request.frame.sid)
|
||||
copy(buf[headerSize:], request.frame.data)
|
||||
|
||||
s.writeLock.Lock()
|
||||
n, err := s.conn.Write(buf[:headerSize+len(request.frame.data)])
|
||||
s.writeLock.Unlock()
|
||||
s.xmitPool.Put(buf)
|
||||
|
||||
n -= headerSize
|
||||
if n < 0 {
|
||||
n = 0
|
||||
}
|
||||
|
||||
result := writeResult{
|
||||
n: n,
|
||||
err: err,
|
||||
}
|
||||
|
||||
request.result <- result
|
||||
close(request.result)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// writeFrame writes the frame to the underlying connection
|
||||
// and returns the number of bytes written if successful
|
||||
func (s *Session) writeFrame(f Frame) (n int, err error) {
|
||||
buf := s.xmitPool.Get().([]byte)
|
||||
buf[0] = f.ver
|
||||
buf[1] = f.cmd
|
||||
binary.LittleEndian.PutUint16(buf[2:], uint16(len(f.data)))
|
||||
binary.LittleEndian.PutUint32(buf[4:], f.sid)
|
||||
copy(buf[headerSize:], f.data)
|
||||
req := writeRequest{
|
||||
frame: f,
|
||||
result: make(chan writeResult, 1),
|
||||
}
|
||||
select {
|
||||
case <-s.die:
|
||||
return 0, errors.New(errBrokenPipe)
|
||||
case s.writes <- req:
|
||||
}
|
||||
|
||||
s.writeLock.Lock()
|
||||
n, err = s.conn.Write(buf[:headerSize+len(f.data)])
|
||||
s.writeLock.Unlock()
|
||||
s.xmitPool.Put(buf)
|
||||
return n, err
|
||||
result := <-req.result
|
||||
return result.n, result.err
|
||||
}
|
||||
|
||||
@@ -500,6 +500,32 @@ func TestReadDeadline(t *testing.T) {
|
||||
session.Close()
|
||||
}
|
||||
|
||||
func TestWriteDeadline(t *testing.T) {
|
||||
cli, err := net.Dial("tcp", "127.0.0.1:19999")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
session, _ := Client(cli, nil)
|
||||
stream, _ := session.OpenStream()
|
||||
const N = 100
|
||||
buf := make([]byte, 10)
|
||||
var writeErr error
|
||||
for i := 0; i < N; i++ {
|
||||
stream.SetWriteDeadline(time.Now().Add(-1 * time.Minute))
|
||||
if _, writeErr = stream.Write(buf); writeErr != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
if writeErr != nil {
|
||||
if !strings.Contains(writeErr.Error(), "i/o timeout") {
|
||||
t.Fatalf("Wrong error: %v", writeErr)
|
||||
}
|
||||
} else {
|
||||
t.Fatal("No error when writing with past deadline")
|
||||
}
|
||||
session.Close()
|
||||
}
|
||||
|
||||
func BenchmarkAcceptClose(b *testing.B) {
|
||||
cli, err := net.Dial("tcp", "127.0.0.1:19999")
|
||||
if err != nil {
|
||||
|
||||
@@ -12,16 +12,17 @@ import (
|
||||
|
||||
// Stream implements io.ReadWriteCloser
|
||||
type Stream struct {
|
||||
id uint32
|
||||
rstflag int32
|
||||
sess *Session
|
||||
buffer bytes.Buffer
|
||||
bufferLock sync.Mutex
|
||||
frameSize int
|
||||
chReadEvent chan struct{} // notify a read event
|
||||
die chan struct{} // flag the stream has closed
|
||||
dieLock sync.Mutex
|
||||
readDeadline atomic.Value
|
||||
id uint32
|
||||
rstflag int32
|
||||
sess *Session
|
||||
buffer bytes.Buffer
|
||||
bufferLock sync.Mutex
|
||||
frameSize int
|
||||
chReadEvent chan struct{} // notify a read event
|
||||
die chan struct{} // flag the stream has closed
|
||||
dieLock sync.Mutex
|
||||
readDeadline atomic.Value
|
||||
writeDeadline atomic.Value
|
||||
}
|
||||
|
||||
// newStream initiates a Stream struct
|
||||
@@ -77,6 +78,13 @@ READ:
|
||||
|
||||
// Write implements io.ReadWriteCloser
|
||||
func (s *Stream) Write(b []byte) (n int, err error) {
|
||||
var deadline <-chan time.Time
|
||||
if d, ok := s.writeDeadline.Load().(time.Time); ok && !d.IsZero() {
|
||||
timer := time.NewTimer(d.Sub(time.Now()))
|
||||
defer timer.Stop()
|
||||
deadline = timer.C
|
||||
}
|
||||
|
||||
select {
|
||||
case <-s.die:
|
||||
return 0, errors.New(errBrokenPipe)
|
||||
@@ -84,12 +92,34 @@ func (s *Stream) Write(b []byte) (n int, err error) {
|
||||
}
|
||||
|
||||
frames := s.split(b, cmdPSH, s.id)
|
||||
sent := 0
|
||||
for k := range frames {
|
||||
if _, err := s.sess.writeFrame(frames[k]); err != nil {
|
||||
return 0, err
|
||||
req := writeRequest{
|
||||
frame: frames[k],
|
||||
result: make(chan writeResult, 1),
|
||||
}
|
||||
|
||||
select {
|
||||
case s.sess.writes <- req:
|
||||
case <-s.die:
|
||||
return sent, errors.New(errBrokenPipe)
|
||||
case <-deadline:
|
||||
return sent, errTimeout
|
||||
}
|
||||
|
||||
select {
|
||||
case result := <-req.result:
|
||||
sent += result.n
|
||||
if result.err != nil {
|
||||
return sent, result.err
|
||||
}
|
||||
case <-s.die:
|
||||
return sent, errors.New(errBrokenPipe)
|
||||
case <-deadline:
|
||||
return sent, errTimeout
|
||||
}
|
||||
}
|
||||
return len(b), nil
|
||||
return sent, nil
|
||||
}
|
||||
|
||||
// Close implements io.ReadWriteCloser
|
||||
@@ -116,6 +146,27 @@ func (s *Stream) SetReadDeadline(t time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetWriteDeadline sets the write deadline as defined by
|
||||
// net.Conn.SetWriteDeadline.
|
||||
// A zero time value disables the deadline.
|
||||
func (s *Stream) SetWriteDeadline(t time.Time) error {
|
||||
s.writeDeadline.Store(t)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetDeadline sets both read and write deadlines as defined by
|
||||
// net.Conn.SetDeadline.
|
||||
// A zero time value disables the deadlines.
|
||||
func (s *Stream) SetDeadline(t time.Time) error {
|
||||
if err := s.SetReadDeadline(t); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.SetWriteDeadline(t); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// session closes the stream
|
||||
func (s *Stream) sessionClose() {
|
||||
s.dieLock.Lock()
|
||||
|
||||
Reference in New Issue
Block a user