From 8510e3505d1b808b7b2b204faca5319d0475c24d Mon Sep 17 00:00:00 2001 From: Audrius Butkevicius Date: Sun, 6 Nov 2016 22:32:44 +0000 Subject: [PATCH] Add WriteDeadline --- session.go | 90 +++++++++++++++++++++++++++++++++++++++++++------ session_test.go | 26 ++++++++++++++ stream.go | 77 +++++++++++++++++++++++++++++++++++------- 3 files changed, 169 insertions(+), 24 deletions(-) diff --git a/session.go b/session.go index 53ed0f5..ca6e651 100644 --- a/session.go +++ b/session.go @@ -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 } diff --git a/session_test.go b/session_test.go index 0b53430..eee0385 100644 --- a/session_test.go +++ b/session_test.go @@ -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 { diff --git a/stream.go b/stream.go index be734a5..6c0368b 100644 --- a/stream.go +++ b/stream.go @@ -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()