Skip to content

Commit a44e87a

Browse files
committed
test conn close/state
1 parent 2d66654 commit a44e87a

2 files changed

Lines changed: 51 additions & 6 deletions

File tree

websocket.go

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -380,8 +380,12 @@ func (ws *Websocket) doCloseHandshake(closeFrame *Frame, cause error) error {
380380
// also ensure no one read or write can exceed our new close deadline,
381381
// since we may do multiple reads below while waiting for our close
382382
// frame to be ACK'd
383-
ws.writeTimeout = ws.closeTimeout
384-
ws.readTimeout = ws.closeTimeout
383+
if ws.readTimeout > 0 {
384+
ws.readTimeout = ws.closeTimeout
385+
}
386+
if ws.writeTimeout > 0 {
387+
ws.writeTimeout = ws.closeTimeout
388+
}
385389

386390
ws.hooks.OnCloseHandshakeStart(ws.clientKey, 0, cause) // TODO: close code
387391
if err := ws.writeFrame(closeFrame); err != nil && !errors.Is(err, net.ErrClosed) {

websocket_test.go

Lines changed: 45 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -747,6 +747,40 @@ func TestNew(t *testing.T) {
747747
})
748748
}
749749

750+
func TestClose(t *testing.T) {
751+
// simulate a websocket client and server connaction
752+
clientConn, serverConn := net.Pipe()
753+
ws := websocket.New(serverConn, websocket.ClientKey("test-client-key"), websocket.ServerMode, websocket.Options{})
754+
755+
// close the server connection
756+
var wg sync.WaitGroup
757+
wg.Add(1)
758+
go func() {
759+
defer wg.Done()
760+
assert.NilError(t, ws.Close())
761+
}()
762+
763+
// confirm that the client got the close frame as expected
764+
wg.Add(1)
765+
go func() {
766+
defer wg.Done()
767+
mustReadCloseFrame(t, clientConn, websocket.StatusNormalClosure, nil)
768+
mustWriteFrame(t, clientConn, true, websocket.NewCloseFrame(websocket.StatusNormalClosure, ""))
769+
}()
770+
wg.Wait()
771+
772+
// make sure any reads or writes after closing the connection are rejected
773+
{
774+
msg, err := ws.ReadMessage(t.Context())
775+
assert.Equal(t, msg, nil, "msg should be nil")
776+
assert.Error(t, err, websocket.ErrConnectionClosed)
777+
}
778+
{
779+
err := ws.WriteMessage(t.Context(), &websocket.Message{})
780+
assert.Error(t, err, websocket.ErrConnectionClosed)
781+
}
782+
}
783+
750784
func mustReadFrame(t testing.TB, src io.Reader, maxPayloadLen int) *websocket.Frame {
751785
t.Helper()
752786
frame, err := websocket.ReadFrame(src, websocket.ClientMode, maxPayloadLen)
@@ -994,27 +1028,34 @@ var (
9941028
// provide a separate reader and writer, and ensures that the conn can't be
9951029
// used after it's closed.
9961030
type dummyConn struct {
1031+
mu sync.Mutex
9971032
in io.Reader
9981033
out io.Writer
999-
closed atomic.Bool
1034+
closed bool
10001035
}
10011036

10021037
func (c *dummyConn) Read(p []byte) (int, error) {
1003-
if c.closed.Load() {
1038+
c.mu.Lock()
1039+
defer c.mu.Unlock()
1040+
if c.closed {
10041041
return 0, errors.New("reader closed")
10051042
}
10061043
return c.in.Read(p)
10071044
}
10081045

10091046
func (c *dummyConn) Write(p []byte) (int, error) {
1010-
if c.closed.Load() {
1047+
c.mu.Lock()
1048+
defer c.mu.Unlock()
1049+
if c.closed {
10111050
return 0, errors.New("writer closed")
10121051
}
10131052
return c.out.Write(p)
10141053
}
10151054

10161055
func (c *dummyConn) Close() error {
1017-
c.closed.Swap(true)
1056+
c.mu.Lock()
1057+
defer c.mu.Unlock()
1058+
c.closed = true
10181059
return nil
10191060
}
10201061

0 commit comments

Comments
 (0)