diff --git a/go.mod b/go.mod index d32fbd77..336411a5 100644 --- a/go.mod +++ b/go.mod @@ -1,3 +1,3 @@ module github.com/coder/websocket -go 1.23 +go 1.19 diff --git a/internal/wsjs/wsjs_js.go b/internal/wsjs/wsjs_js.go index 45ecf49d..49895e33 100644 --- a/internal/wsjs/wsjs_js.go +++ b/internal/wsjs/wsjs_js.go @@ -139,6 +139,13 @@ func (c WebSocket) Close(code int, reason string) (err error) { return err } +// CloseNow closes the WebSocket without a status code or reason. +func (c WebSocket) CloseNow() (err error) { + defer handleJSError(&err, nil) + c.v.Call("close") + return err +} + // SendText sends the given string as a text message // on the WebSocket. func (c WebSocket) SendText(v string) (err error) { @@ -155,6 +162,15 @@ func (c WebSocket) SendBytes(v []byte) (err error) { return err } +// BufferedAmount returns the number of bytes of data that have been queued +// using calls to send() but not yet transmitted to the network. +// This value resets to zero once all queued data has been sent. +// Use this to implement backpressure detection - if this value grows over time, +// the client is falling behind and may need throttling or reconnection. +func (c WebSocket) BufferedAmount() uint64 { + return uint64(c.v.Get("bufferedAmount").Int()) +} + func extractArrayBuffer(arrayBuffer js.Value) []byte { uint8Array := js.Global().Get("Uint8Array").New(arrayBuffer) dst := make([]byte, uint8Array.Length()) diff --git a/ws_js.go b/ws_js.go index 026b75fc..6205a569 100644 --- a/ws_js.go +++ b/ws_js.go @@ -57,6 +57,7 @@ type Conn struct { closeErr error closeWasClean bool + releaseOnce sync.Once releaseOnClose func() releaseOnError func() releaseOnMessage func() @@ -94,15 +95,11 @@ func (c *Conn) init() { // its possible the browser triggered it without us // explicitly sending it. c.close(err, e.WasClean) - - c.releaseOnClose() - c.releaseOnError() - c.releaseOnMessage() + c.releaseEventHandlers() }) c.releaseOnError = c.ws.OnError(func(v js.Value) { - c.setCloseErr(errors.New(v.Get("message").String())) - c.closeWithInternal() + c.abort(errors.New(v.Get("message").String())) }) c.releaseOnMessage = c.ws.OnMessage(func(e wsjs.MessageEvent) { @@ -124,8 +121,22 @@ func (c *Conn) init() { }) } +func (c *Conn) releaseEventHandlers() { + c.releaseOnce.Do(func() { + c.releaseOnClose() + c.releaseOnError() + c.releaseOnMessage() + }) +} + +func (c *Conn) abort(err error) { + c.close(err, false) + c.releaseEventHandlers() + _ = c.ws.CloseNow() +} + func (c *Conn) closeWithInternal() { - c.Close(StatusInternalError, "something went wrong") + c.abort(errors.New("something went wrong")) } // Read attempts to read a message from the connection. @@ -238,11 +249,12 @@ func (c *Conn) Close(code StatusCode, reason string) error { // CloseNow closes the WebSocket connection without attempting a close handshake. // Use when you do not want the overhead of the close handshake. -// -// note: No different from Close(StatusGoingAway, "") in WASM as there is no way to close -// a WebSocket without the close handshake. func (c *Conn) CloseNow() error { - return c.Close(StatusGoingAway, "") + if c.isClosed() { + return net.ErrClosed + } + c.abort(fmt.Errorf("sent close: %w", CloseError{Code: StatusGoingAway})) + return nil } func (c *Conn) exportedClose(code StatusCode, reason string) error { @@ -261,6 +273,7 @@ func (c *Conn) exportedClose(code StatusCode, reason string) error { c.setCloseErr(ce) err := c.ws.Close(int(code), reason) if err != nil { + c.abort(ce) return err } @@ -321,7 +334,7 @@ func dial(ctx context.Context, url string, opts *DialOptions) (*Conn, *http.Resp select { case <-ctx.Done(): - c.Close(StatusPolicyViolation, "dial timed out") + c.abort(ctx.Err()) return nil, nil, ctx.Err() case <-opench: return c, &http.Response{ @@ -417,6 +430,14 @@ func (c *Conn) SetReadLimit(n int64) { c.msgReadLimit.Store(n) } +// BufferedAmount returns the number of bytes of data that have been queued +// using calls to Write() but not yet transmitted to the network. +// Monitor this to detect backpressure - a growing value indicates the browser +// is falling behind and may need throttling or reconnection. +func (c *Conn) BufferedAmount() uint64 { + return c.ws.BufferedAmount() +} + func (c *Conn) setCloseErr(err error) { c.closeErrOnce.Do(func() { c.closeErr = fmt.Errorf("WebSocket closed: %w", err) diff --git a/ws_js_test.go b/ws_js_test.go index 1fa242f4..80d100ae 100644 --- a/ws_js_test.go +++ b/ws_js_test.go @@ -4,6 +4,7 @@ import ( "context" "net/http" "os" + "syscall/js" "testing" "time" @@ -38,16 +39,38 @@ func TestWasm(t *testing.T) { } func TestWasmDialTimeout(t *testing.T) { - t.Parallel() - - ctx, cancel := context.WithTimeout(context.Background(), time.Millisecond) - defer cancel() + js.Global().Call("eval", `(() => { + globalThis.__originalWebSocket = globalThis.WebSocket; + globalThis.__activeWebSocketListeners = 0; + const trackedEvents = new Set(["close", "error", "message", "open"]); + globalThis.WebSocket = class extends globalThis.__originalWebSocket { + addEventListener(type, listener, options) { + if (trackedEvents.has(type)) globalThis.__activeWebSocketListeners++; + return super.addEventListener(type, listener, options); + } + removeEventListener(type, listener, options) { + if (trackedEvents.has(type)) globalThis.__activeWebSocketListeners--; + return super.removeEventListener(type, listener, options); + } + }; + })()`) + defer js.Global().Call("eval", `(() => { + globalThis.WebSocket = globalThis.__originalWebSocket; + delete globalThis.__originalWebSocket; + delete globalThis.__activeWebSocketListeners; + })()`) beforeDial := time.Now() - _, _, err := websocket.Dial(ctx, "ws://example.com:9893", &websocket.DialOptions{ - Subprotocols: []string{"echo"}, - }) - assert.Error(t, err) + for i := 0; i < 10; i++ { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, _, err := websocket.Dial(ctx, "ws://example.com:9893", &websocket.DialOptions{ + Subprotocols: []string{"echo"}, + }) + assert.Error(t, err) + assert.Equal(t, "active WebSocket event listeners", 0, js.Global().Get("__activeWebSocketListeners").Int()) + } if time.Since(beforeDial) >= time.Second { t.Fatal("wasm context dial timeout is not working", time.Since(beforeDial)) }