@@ -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+
750784func 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.
9961030type dummyConn struct {
1031+ mu sync.Mutex
9971032 in io.Reader
9981033 out io.Writer
999- closed atomic. Bool
1034+ closed bool
10001035}
10011036
10021037func (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
10091046func (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
10161055func (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