Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 6 additions & 5 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ go get github.com/coder/websocket
- [net.Conn](https://pkg.go.dev/github.com/coder/websocket#NetConn) wrapper
- [Ping pong](https://pkg.go.dev/github.com/coder/websocket#Conn.Ping) API
- [RFC 7692](https://tools.ietf.org/html/rfc7692) permessage-deflate compression
- [Prepared message](https://pkg.go.dev/github.com/coder/websocket#PreparedMessage) for efficient broadcasting
- [CloseRead](https://pkg.go.dev/github.com/coder/websocket#Conn.CloseRead) helper for write only connections
- Compile to [Wasm](https://pkg.go.dev/github.com/coder/websocket#hdr-Wasm)

Expand All @@ -48,10 +49,11 @@ See GitHub issues for minor issues but the major future enhancements are:

## Examples

For a production quality example that demonstrates the complete API, see the
[echo example](./internal/examples/echo).

For a full stack example, see the [chat example](./internal/examples/chat).
- Production quality example that demonstrates the complete API, see the
[echo example](./internal/examples/echo).
- Full stack example, see the [chat example](./internal/examples/chat).
- Broadcasting a message to many connections, see the
[broadcast example](./internal/examples/broadcast).

### Server

Expand Down Expand Up @@ -107,7 +109,6 @@ c.Close(websocket.StatusNormalClosure, "")
Advantages of [gorilla/websocket](https://github.com/gorilla/websocket):

- Mature and widely used
- [Prepared writes](https://pkg.go.dev/github.com/gorilla/websocket#PreparedMessage)
- Configurable [buffer sizes](https://pkg.go.dev/github.com/gorilla/websocket#hdr-Buffers)

Advantages of github.com/coder/websocket:
Expand Down
5 changes: 4 additions & 1 deletion compress.go
Original file line number Diff line number Diff line change
Expand Up @@ -41,10 +41,13 @@ const (
//
// This means less efficient compression as the sliding window from previous messages will not be used but the
// memory overhead will be lower as there will be no fixed cost for the flate.Writer nor the 32 KB sliding window.
// Especially if the connections are long lived and seldom written to.
// Especially if the connections are long-lived and seldom written to.
//
// Thus, it uses less memory than CompressionContextTakeover but compresses less efficiently.
//
// A message written with Conn.Write or Conn.WritePrepared is sent uncompressed if compression does not make it
// smaller, such as already compressed data.
//
// If the peer does not support CompressionNoContextTakeover then we will fall back to CompressionDisabled.
CompressionNoContextTakeover
)
Expand Down
9 changes: 9 additions & 0 deletions compress_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,8 @@ func TestWriteSingleFrameCompressed(t *testing.T) {

largeMsg = []byte(strings.Repeat("hello world ", 100))
smallMsg = []byte("small message")
// Random bytes do not compress, so compression makes them larger.
incompressibleMsg = xrand.Bytes(1024)
)

testCases := []struct {
Expand All @@ -88,6 +90,10 @@ func TestWriteSingleFrameCompressed(t *testing.T) {
{"NoContextTakeover/AboveThreshold", CompressionNoContextTakeover, largeMsg, true},
{"ContextTakeover/BelowThreshold", CompressionContextTakeover, smallMsg, false},
{"NoContextTakeover/BelowThreshold", CompressionNoContextTakeover, smallMsg, false},
// With context takeover, the message is in the sliding window, so it
// stays compressed to keep the peer's window in sync.
{"ContextTakeover/Incompressible", CompressionContextTakeover, incompressibleMsg, true},
{"NoContextTakeover/Incompressible", CompressionNoContextTakeover, incompressibleMsg, false},
}

for _, tc := range testCases {
Expand Down Expand Up @@ -126,6 +132,9 @@ func TestWriteSingleFrameCompressed(t *testing.T) {

assert.Equal(t, "opcode", opText, h.opcode)
assert.Equal(t, "rsv1 (compressed)", tc.wantRsv1, h.rsv1)
if !tc.wantRsv1 {
assert.Equal(t, "payload length", int64(len(tc.msg)), h.payloadLength)
}
assert.Equal(t, "fin", true, h.fin)

err = <-writeDone
Expand Down
147 changes: 147 additions & 0 deletions conn_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -50,9 +50,14 @@ func TestConn(t *testing.T) {

c1.SetReadLimit(131072)

// Interleave prepared writes with regular writes, as they share
// the compression state of the connection.
for range 5 {
err := wstest.Echo(tt.ctx, c1, 131072)
assert.Success(t, err)

err = wstest.EchoPrepared(tt.ctx, c1, 131072)
assert.Success(t, err)
}

err := c1.Close(websocket.StatusNormalClosure, "")
Expand All @@ -61,6 +66,69 @@ func TestConn(t *testing.T) {
}
})

t.Run("writePreparedConcurrent", func(t *testing.T) {
t.Parallel()

ctx, cancel := context.WithTimeout(context.Background(), time.Second*30)
defer cancel()

exp := strings.Repeat("prepared", 128)
pm := websocket.NewPreparedMessage(websocket.MessageText, []byte(exp))

modes := []websocket.CompressionMode{
websocket.CompressionDisabled,
websocket.CompressionContextTakeover,
websocket.CompressionNoContextTakeover,
}

// Write the same message with every compression mode from both sides,
// so connections with different settings share it and client masking
// runs concurrently on the shared payload.
var errs []<-chan error
for i := range 12 {
mode := modes[i%len(modes)]
client, server := wstest.Pipe(&websocket.DialOptions{
CompressionMode: mode,
}, &websocket.AcceptOptions{
CompressionMode: mode,
})
t.Cleanup(func() {
client.CloseNow()
server.CloseNow()
})

w, r := client, server
if i%2 == 0 {
w, r = server, client
}
errs = append(errs, xsync.Go(func() error {
for range 5 {
err := w.WritePrepared(ctx, pm)
if err != nil {
return err
}
}
return nil
}))
errs = append(errs, xsync.Go(func() error {
for range 5 {
_, p, err := r.Read(ctx)
if err != nil {
return err
}
if string(p) != exp {
return fmt.Errorf("unexpected msg: %q", p)
}
}
return nil
}))
}

for _, errc := range errs {
assert.Success(t, <-errc)
}
})

t.Run("badClose", func(t *testing.T) {
tt, c1, c2 := newConnTest(t, nil, nil)

Expand Down Expand Up @@ -667,6 +735,85 @@ func BenchmarkConn(b *testing.B) {
}
}

func BenchmarkBroadcast(b *testing.B) {
const conns = 16

modes := []struct {
name string
mode websocket.CompressionMode
}{
{"disabledCompress", websocket.CompressionDisabled},
{"compressContextTakeover", websocket.CompressionContextTakeover},
{"compressNoContext", websocket.CompressionNoContextTakeover},
}
writes := []struct {
name string
write func(ctx context.Context, cs []*websocket.Conn, msg []byte) error
}{
{"write", func(ctx context.Context, cs []*websocket.Conn, msg []byte) error {
for _, c := range cs {
err := c.Write(ctx, websocket.MessageText, msg)
if err != nil {
return err
}
}
return nil
}},
{"writePrepared", func(ctx context.Context, cs []*websocket.Conn, msg []byte) error {
pm := websocket.NewPreparedMessage(websocket.MessageText, msg)
for _, c := range cs {
err := c.WritePrepared(ctx, pm)
if err != nil {
return err
}
}
return nil
}},
}

msg := []byte(strings.Repeat("1234", 256))

// Clients mask every write, so they are benchmarked separately.
for _, side := range []string{"server", "client"} {
for _, m := range modes {
for _, w := range writes {
b.Run(side+"/"+m.name+"/"+w.name, func(b *testing.B) {
ctx, cancel := context.WithTimeout(context.Background(), time.Minute)
defer cancel()

writers := make([]*websocket.Conn, conns)
for i := range writers {
client, server := wstest.Pipe(&websocket.DialOptions{
CompressionMode: m.mode,
}, &websocket.AcceptOptions{
CompressionMode: m.mode,
})
b.Cleanup(func() {
client.CloseNow()
server.CloseNow()
})
writers[i] = server
if side == "client" {
writers[i] = client
}
writers[i].DiscardWrites()
}

b.SetBytes(int64(len(msg) * conns))
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
err := w.write(ctx, writers, msg)
if err != nil {
b.Fatal(err)
}
}
})
}
}
}
}

func echoServer(w http.ResponseWriter, r *http.Request, opts *websocket.AcceptOptions) (err error) {
defer errd.Wrap(&err, "echo server failed")

Expand Down
10 changes: 8 additions & 2 deletions example_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -118,7 +118,7 @@ func Example_writeOnly() {
}

func Example_crossOrigin() {
// This handler demonstrates how to safely accept cross origin WebSockets
// This handler demonstrates how to safely accept cross-origin WebSockets
// from the origin example.com.
fn := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
c, err := websocket.Accept(w, r, &websocket.AcceptOptions{
Expand Down Expand Up @@ -165,7 +165,13 @@ func Example_fullStackChat() {
// https://github.com/nhooyr/websocket/tree/master/internal/examples/chat
}

// This example demonstrates a echo server.
// This example demonstrates an echo server.
func Example_echo() {
// https://github.com/nhooyr/websocket/tree/master/internal/examples/echo
}

// This example demonstrates broadcasting messages to many subscribers with
// a PreparedMessage.
func ExampleConn_WritePrepared() {
// https://github.com/coder/websocket/tree/master/internal/examples/broadcast
}
5 changes: 5 additions & 0 deletions export_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
package websocket

import (
"io"
"net"

"github.com/coder/websocket/internal/util"
Expand All @@ -17,6 +18,10 @@ func (c *Conn) RecordBytesWritten() *int {
return &bytesWritten
}

func (c *Conn) DiscardWrites() {
c.bw.Reset(io.Discard)
}

func (c *Conn) RecordBytesRead() *int {
var bytesRead int
c.br.Reset(util.ReaderFunc(func(p []byte) (int, error) {
Expand Down
38 changes: 38 additions & 0 deletions internal/examples/broadcast/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
# Broadcast Example

This directory contains a broadcast server example using github.com/coder/websocket.

```bash
$ cd internal/examples/broadcast
$ go run . localhost:8080
listening on http://127.0.0.1:8080
```

Subscribe with a WebSocket client like [websocat](https://github.com/vi/websocat) and publish
messages with curl:

```bash
$ websocat ws://127.0.0.1:8080/subscribe
$ curl --data-binary 'hello' http://127.0.0.1:8080/publish
```

Every published message is delivered to all subscribers. Messages are sent as text, so
`/publish` rejects bodies that are not valid UTF-8.

## Structure

The server is in `server.go`. Subscribers connect to the WebSocket `/subscribe` endpoint and
messages are published via the HTTP POST `/publish` endpoint, so that you can easily publish
with curl.

Each published message is wrapped in a single `PreparedMessage` and queued for every subscriber.
Connections are accepted with `CompressionNoContextTakeover`, so a message that reaches the
compression threshold is compressed once per broadcast instead of once per subscriber.
Subscribers that do not support compression receive the same message uncompressed.

Publishing never blocks: a subscriber that cannot keep up with its queue is disconnected.

`server_test.go` contains a test that subscribes clients with and without compression and
ensures every client receives every published message, while invalid UTF-8 is rejected.

`main.go` brings it all together so that you can run it and play around with it.
59 changes: 59 additions & 0 deletions internal/examples/broadcast/main.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
package main

import (
"context"
"errors"
"log"
"net"
"net/http"
"os"
"os/signal"
"time"
)

func main() {
log.SetFlags(0)

err := run()
if err != nil {
log.Fatal(err)
}
}

// run starts a http.Server for the passed in address
// with all requests handled by broadcastServer.
func run() error {
if len(os.Args) < 2 {
return errors.New("please provide an address to listen on as the first argument")
}

l, err := net.Listen("tcp", os.Args[1])
if err != nil {
return err
}
log.Printf("listening on http://%v", l.Addr())

s := &http.Server{
Handler: newBroadcastServer(log.Printf),
ReadTimeout: time.Second * 10,
WriteTimeout: time.Second * 10,
}
errc := make(chan error, 1)
go func() {
errc <- s.Serve(l)
}()

sigs := make(chan os.Signal, 1)
signal.Notify(sigs, os.Interrupt)
select {
case err := <-errc:
log.Printf("failed to serve: %v", err)
case sig := <-sigs:
log.Printf("terminating: %v", sig)
}

ctx, cancel := context.WithTimeout(context.Background(), time.Second*10)
defer cancel()

return s.Shutdown(ctx)
}
Loading
Loading