diff --git a/session.go b/session.go index 12fc4cb..23671da 100644 --- a/session.go +++ b/session.go @@ -118,7 +118,7 @@ func (s *Session) OpenStream() (*Stream, error) { 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())) + timer := time.NewTimer(time.Until(d)) defer timer.Stop() deadline = timer.C } diff --git a/session_test.go b/session_test.go index 2147d8f..760642d 100644 --- a/session_test.go +++ b/session_test.go @@ -5,36 +5,37 @@ import ( "encoding/binary" "fmt" "io" - "log" "math/rand" "net" - "net/http" - _ "net/http/pprof" "strings" "sync" "testing" "time" ) -func init() { - go func() { - log.Println(http.ListenAndServe("localhost:6060", nil)) - }() - log.SetFlags(log.LstdFlags | log.Lshortfile) - ln, err := net.Listen("tcp", "127.0.0.1:19999") +// setupServer starts new server listening on a random localhost port and +// returns address of the server, function to stop the server, new client +// connection to this server or an error. +func setupServer(tb testing.TB) (addr string, stopfunc func(), client net.Conn, err error) { + ln, err := net.Listen("tcp", "localhost:0") if err != nil { - // handle error - panic(err) + return "", nil, nil, err } go func() { - for { - conn, err := ln.Accept() - if err != nil { - // handle error - } - go handleConnection(conn) + conn, err := ln.Accept() + if err != nil { + tb.Error(err) + return } + go handleConnection(conn) }() + addr = ln.Addr().String() + conn, err := net.Dial("tcp", addr) + if err != nil { + ln.Close() + return "", nil, nil, err + } + return ln.Addr().String(), func() { ln.Close() }, conn, nil } func handleConnection(conn net.Conn) { @@ -58,10 +59,11 @@ func handleConnection(conn net.Conn) { } func TestEcho(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() const N = 100 @@ -85,10 +87,11 @@ func TestEcho(t *testing.T) { } func TestSpeed(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() t.Log(stream.LocalAddr(), stream.RemoteAddr()) @@ -102,7 +105,7 @@ func TestSpeed(t *testing.T) { for { n, err := stream.Read(buf) if err != nil { - t.Fatal(err) + t.Error(err) break } else { nrecv += n @@ -124,10 +127,11 @@ func TestSpeed(t *testing.T) { } func TestParallel(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) par := 1000 @@ -155,10 +159,11 @@ func TestParallel(t *testing.T) { } func TestCloseThenOpen(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) session.Close() if _, err := session.OpenStream(); err == nil { @@ -167,10 +172,11 @@ func TestCloseThenOpen(t *testing.T) { } func TestStreamDoubleClose(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() stream.Close() @@ -181,10 +187,11 @@ func TestStreamDoubleClose(t *testing.T) { } func TestConcurrentClose(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) numStreams := 100 streams := make([]*Stream, 0, numStreams) @@ -206,10 +213,11 @@ func TestConcurrentClose(t *testing.T) { } func TestTinyReadBuffer(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() const N = 100 @@ -241,76 +249,91 @@ func TestTinyReadBuffer(t *testing.T) { } func TestIsClose(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) session.Close() - if session.IsClosed() != true { + if !session.IsClosed() { t.Fatal("still open after close") } } func TestKeepAliveTimeout(t *testing.T) { - ln, err := net.Listen("tcp", "127.0.0.1:29999") + ln, err := net.Listen("tcp", "localhost:0") if err != nil { - // handle error - panic(err) + t.Fatal(err) } + defer ln.Close() go func() { ln.Accept() }() - cli, err := net.Dial("tcp", "127.0.0.1:29999") + cli, err := net.Dial("tcp", ln.Addr().String()) if err != nil { t.Fatal(err) } + defer cli.Close() config := DefaultConfig() config.KeepAliveInterval = time.Second config.KeepAliveTimeout = 2 * time.Second session, _ := Client(cli, config) - <-time.After(3 * time.Second) - if session.IsClosed() != true { + time.Sleep(3 * time.Second) + if !session.IsClosed() { t.Fatal("keepalive-timeout failed") } } func TestServerEcho(t *testing.T) { - ln, err := net.Listen("tcp", "127.0.0.1:39999") - if err != nil { - // handle error - panic(err) - } - go func() { - if conn, err := ln.Accept(); err == nil { - session, _ := Server(conn, nil) - if stream, err := session.OpenStream(); err == nil { - const N = 100 - buf := make([]byte, 10) - for i := 0; i < N; i++ { - msg := fmt.Sprintf("hello%v", i) - stream.Write([]byte(msg)) - if n, err := stream.Read(buf); err != nil { - t.Fatal(err) - } else if string(buf[:n]) != msg { - t.Fatal(err) - } - } - stream.Close() - } else { - t.Fatal(err) - } - } else { - t.Fatal(err) - } - }() - - cli, err := net.Dial("tcp", "127.0.0.1:39999") + ln, err := net.Listen("tcp", "localhost:0") if err != nil { t.Fatal(err) } + defer ln.Close() + go func() { + err := func() error { + conn, err := ln.Accept() + if err != nil { + return err + } + defer conn.Close() + session, err := Server(conn, nil) + if err != nil { + return err + } + defer session.Close() + buf := make([]byte, 10) + stream, err := session.OpenStream() + if err != nil { + return err + } + defer stream.Close() + for i := 0; i < 100; i++ { + msg := fmt.Sprintf("hello%v", i) + stream.Write([]byte(msg)) + n, err := stream.Read(buf) + if err != nil { + return err + } + if got := string(buf[:n]); got != msg { + return fmt.Errorf("got: %q, want: %q", got, msg) + } + } + return nil + }() + if err != nil { + t.Error(err) + } + }() + + cli, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + t.Fatal(err) + } + defer cli.Close() if session, err := Client(cli, nil); err == nil { if stream, err := session.AcceptStream(); err == nil { buf := make([]byte, 65536) @@ -330,10 +353,11 @@ func TestServerEcho(t *testing.T) { } func TestSendWithoutRecv(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() const N = 100 @@ -349,10 +373,11 @@ func TestSendWithoutRecv(t *testing.T) { } func TestWriteAfterClose(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() stream.Close() @@ -362,10 +387,11 @@ func TestWriteAfterClose(t *testing.T) { } func TestReadStreamAfterSessionClose(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() session.Close() @@ -378,10 +404,11 @@ func TestReadStreamAfterSessionClose(t *testing.T) { } func TestWriteStreamAfterConnectionClose(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() session.conn.Close() @@ -391,10 +418,11 @@ func TestWriteStreamAfterConnectionClose(t *testing.T) { } func TestNumStreamAfterClose(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) if _, err := session.OpenStream(); err == nil { if session.NumStreams() != 1 { @@ -411,11 +439,12 @@ func TestNumStreamAfterClose(t *testing.T) { } func TestRandomFrame(t *testing.T) { - // pure random - cli, err := net.Dial("tcp", "127.0.0.1:19999") + addr, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() + // pure random session, _ := Client(cli, nil) for i := 0; i < 100; i++ { rnd := make([]byte, rand.Uint32()%1024) @@ -425,7 +454,7 @@ func TestRandomFrame(t *testing.T) { cli.Close() // double syn - cli, err = net.Dial("tcp", "127.0.0.1:19999") + cli, err = net.Dial("tcp", addr) if err != nil { t.Fatal(err) } @@ -437,7 +466,7 @@ func TestRandomFrame(t *testing.T) { cli.Close() // random cmds - cli, err = net.Dial("tcp", "127.0.0.1:19999") + cli, err = net.Dial("tcp", addr) if err != nil { t.Fatal(err) } @@ -450,7 +479,7 @@ func TestRandomFrame(t *testing.T) { cli.Close() // random cmds & sids - cli, err = net.Dial("tcp", "127.0.0.1:19999") + cli, err = net.Dial("tcp", addr) if err != nil { t.Fatal(err) } @@ -462,7 +491,7 @@ func TestRandomFrame(t *testing.T) { cli.Close() // random version - cli, err = net.Dial("tcp", "127.0.0.1:19999") + cli, err = net.Dial("tcp", addr) if err != nil { t.Fatal(err) } @@ -475,7 +504,7 @@ func TestRandomFrame(t *testing.T) { cli.Close() // incorrect size - cli, err = net.Dial("tcp", "127.0.0.1:19999") + cli, err = net.Dial("tcp", addr) if err != nil { t.Fatal(err) } @@ -499,10 +528,11 @@ func TestRandomFrame(t *testing.T) { } func TestReadDeadline(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() const N = 100 @@ -527,10 +557,11 @@ func TestReadDeadline(t *testing.T) { } func TestWriteDeadline(t *testing.T) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(t) if err != nil { t.Fatal(err) } + defer stop() session, _ := Client(cli, nil) stream, _ := session.OpenStream() buf := make([]byte, 10) @@ -548,10 +579,11 @@ func TestWriteDeadline(t *testing.T) { } func BenchmarkAcceptClose(b *testing.B) { - cli, err := net.Dial("tcp", "127.0.0.1:19999") + _, stop, cli, err := setupServer(b) if err != nil { b.Fatal(err) } + defer stop() session, _ := Client(cli, nil) for i := 0; i < b.N; i++ { if stream, err := session.OpenStream(); err == nil { @@ -616,10 +648,11 @@ func getSmuxStreamPair() (*Stream, *Stream, error) { } func getTCPConnectionPair() (net.Conn, net.Conn, error) { - lst, err := net.Listen("tcp", "127.0.0.1:0") + lst, err := net.Listen("tcp", "localhost:0") if err != nil { return nil, nil, err } + defer lst.Close() var conn0 net.Conn var err0 error diff --git a/stream.go b/stream.go index 613bd63..1b3ebe0 100644 --- a/stream.go +++ b/stream.go @@ -46,7 +46,7 @@ func (s *Stream) ID() uint32 { func (s *Stream) Read(b []byte) (n int, err error) { var deadline <-chan time.Time if d, ok := s.readDeadline.Load().(time.Time); ok && !d.IsZero() { - timer := time.NewTimer(d.Sub(time.Now())) + timer := time.NewTimer(time.Until(d)) defer timer.Stop() deadline = timer.C } @@ -78,7 +78,7 @@ READ: 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())) + timer := time.NewTimer(time.Until(d)) defer timer.Stop() deadline = timer.C }