Skip to content
Merged
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
9 changes: 4 additions & 5 deletions autobahn_test.go
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
package websocket_test
package websocket

// ============================================================================
// Autobahn Test Suite
Expand Down Expand Up @@ -41,7 +41,6 @@ import (
"testing"
"time"

"github.com/mccutchen/websocket"
"github.com/mccutchen/websocket/internal/testing/assert"
)

Expand Down Expand Up @@ -74,15 +73,15 @@ func TestAutobahn(t *testing.T) {
}

// Hooks can be expensive, so only enable them if necessary for debugging
var hooks websocket.Hooks
var hooks Hooks
if debug := os.Getenv("DEBUG"); debug == "1" {
hooks = newTestHooks(t)
}

targetURL := os.Getenv("TARGET")
if targetURL == "" {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ws, err := websocket.Accept(w, r, websocket.Options{
ws, err := Accept(w, r, Options{
Hooks: hooks,
// long ReadTimeout because some autobahn test cases (e.g. 5.19)
// sleep up to 1 second between frames
Expand All @@ -97,7 +96,7 @@ func TestAutobahn(t *testing.T) {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
_ = ws.Handle(r.Context(), websocket.EchoHandler)
_ = ws.Handle(r.Context(), EchoHandler)
}))
defer srv.Close()
targetURL = srv.URL
Expand Down
73 changes: 36 additions & 37 deletions proto_test.go
Original file line number Diff line number Diff line change
@@ -1,11 +1,10 @@
package websocket_test
package websocket

import (
"bytes"
"fmt"
"testing"

"github.com/mccutchen/websocket"
"github.com/mccutchen/websocket/internal/testing/assert"
)

Expand All @@ -15,12 +14,12 @@ func TestFrameRoundTrip(t *testing.T) {
t.Parallel()

// write masked "client" frame to buffer
clientFrame := websocket.NewFrame(websocket.OpcodeText, true, []byte("hello"))
clientFrame := NewFrame(OpcodeText, true, []byte("hello"))
buf := &bytes.Buffer{}
assert.NilError(t, websocket.WriteFrame(buf, websocket.NewMaskingKey(), clientFrame))
assert.NilError(t, WriteFrame(buf, NewMaskingKey(), clientFrame))

// read "server" frame from buffer.
serverFrame, err := websocket.ReadFrame(buf, websocket.ServerMode, len(clientFrame.Payload))
serverFrame, err := ReadFrame(buf, ServerMode, len(clientFrame.Payload))
assert.NilError(t, err)

// ensure client and server frame match
Expand All @@ -33,23 +32,23 @@ func TestMaxFrameSize(t *testing.T) {
t.Parallel()

// write masked "client" frame to buffer
clientFrame := websocket.NewFrame(websocket.OpcodeText, true, []byte("hello"))
clientFrame := NewFrame(OpcodeText, true, []byte("hello"))
buf := &bytes.Buffer{}
assert.NilError(t, websocket.WriteFrame(buf, websocket.NewMaskingKey(), clientFrame))
assert.NilError(t, WriteFrame(buf, NewMaskingKey(), clientFrame))

// read "server" frame from buffer.
serverFrame, err := websocket.ReadFrame(buf, websocket.ServerMode, len(clientFrame.Payload)-1)
assert.Error(t, err, websocket.ErrFrameTooLarge)
serverFrame, err := ReadFrame(buf, ServerMode, len(clientFrame.Payload)-1)
assert.Error(t, err, ErrFrameTooLarge)
assert.Equal(t, serverFrame, nil, "expected nil frame on error")
}

func TestRSV(t *testing.T) {
// We don't currently support any extensions, so RSV bits are not allowed.
// But we still need to be able properly parse and marshal them.
marshalledFrame := func(rsvBits ...websocket.RSVBit) []byte {
marshalledFrame := func(rsvBits ...RSVBit) []byte {
buf := &bytes.Buffer{}
frame := websocket.NewFrame(websocket.OpcodeText, true, nil, rsvBits...)
assert.NilError(t, websocket.WriteFrame(buf, websocket.Unmasked, frame))
frame := NewFrame(OpcodeText, true, nil, rsvBits...)
assert.NilError(t, WriteFrame(buf, Unmasked, frame))
return buf.Bytes()
}

Expand All @@ -63,19 +62,19 @@ func TestRSV(t *testing.T) {
rawBytes: marshalledFrame(),
},
"RSV1 set": {
rawBytes: marshalledFrame(websocket.RSV1),
rawBytes: marshalledFrame(RSV1),
wantRSV1: true,
},
"RSV2 set": {
rawBytes: marshalledFrame(websocket.RSV2),
rawBytes: marshalledFrame(RSV2),
wantRSV2: true,
},
"RSV3 set": {
rawBytes: marshalledFrame(websocket.RSV3),
rawBytes: marshalledFrame(RSV3),
wantRSV3: true,
},
"all RSV bits set": {
rawBytes: marshalledFrame(websocket.RSV1, websocket.RSV2, websocket.RSV3),
rawBytes: marshalledFrame(RSV1, RSV2, RSV3),
wantRSV1: true,
wantRSV2: true,
wantRSV3: true,
Expand All @@ -98,47 +97,47 @@ func TestExampleFramesFromRFC(t *testing.T) {
// https://datatracker.ietf.org/doc/html/rfc6455#section-5.7
testCases := map[string]struct {
rawBytes []byte
wantFrame *websocket.Frame
wantFrame *Frame
}{
"single-frame unmasked text": {
rawBytes: []byte{0x81, 0x05, 0x48, 0x65, 0x6c, 0x6c, 0x6f},
wantFrame: websocket.NewFrame(websocket.OpcodeText, true, []byte("Hello")),
wantFrame: NewFrame(OpcodeText, true, []byte("Hello")),
},
"single-frame masked text": {
rawBytes: []byte{0x81, 0x85, 0x37, 0xfa, 0x21, 0x3d, 0x7f, 0x9f, 0x4d, 0x51, 0x58},
wantFrame: websocket.NewFrame(websocket.OpcodeText, true, []byte("Hello")),
wantFrame: NewFrame(OpcodeText, true, []byte("Hello")),
},
"fragmented unmasked text part 1": {
rawBytes: []byte{0x01, 0x03, 0x48, 0x65, 0x6c},
wantFrame: websocket.NewFrame(websocket.OpcodeText, false, []byte("Hel")),
wantFrame: NewFrame(OpcodeText, false, []byte("Hel")),
},
"fragmented unmasked text part 2": {
rawBytes: []byte{0x80, 0x02, 0x6c, 0x6f},
wantFrame: websocket.NewFrame(websocket.OpcodeContinuation, true, []byte("lo")),
wantFrame: NewFrame(OpcodeContinuation, true, []byte("lo")),
},
"unmasked ping": {
rawBytes: []byte{
0x89, 0x05, 0x48, 0x65, 0x6c, 0x6c, 0x6f,
},
wantFrame: websocket.NewFrame(websocket.OpcodePing, true, []byte("Hello")),
wantFrame: NewFrame(OpcodePing, true, []byte("Hello")),
},
"masked ping response": {
rawBytes: []byte{0x8a, 0x85, 0x37, 0xfa, 0x21, 0x3d, 0x7f, 0x9f, 0x4d, 0x51, 0x58},
wantFrame: websocket.NewFrame(websocket.OpcodePong, true, []byte("Hello")),
wantFrame: NewFrame(OpcodePong, true, []byte("Hello")),
},
"256 bytes binary message": {
rawBytes: append(
[]byte{0x82, 0x7E, 0x01, 0x00},
make([]byte, 256)...,
),
wantFrame: websocket.NewFrame(websocket.OpcodeBinary, true, make([]byte, 256)),
wantFrame: NewFrame(OpcodeBinary, true, make([]byte, 256)),
},
"64KiB binary message": {
rawBytes: append(
[]byte{0x82, 0x7F, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00},
make([]byte, 65536)...,
),
wantFrame: websocket.NewFrame(websocket.OpcodeBinary, true, make([]byte, 65536)),
wantFrame: NewFrame(OpcodeBinary, true, make([]byte, 65536)),
},
}

Expand Down Expand Up @@ -179,7 +178,7 @@ func TestIncompleteFrames(t *testing.T) {
t.Run(name, func(t *testing.T) {
t.Parallel()
buf := bytes.NewReader(tc.rawBytes)
_, err := websocket.ReadFrame(buf, websocket.ClientMode, 70000)
_, err := ReadFrame(buf, ClientMode, 70000)
assert.Error(t, err, tc.wantErr)
})
}
Expand Down Expand Up @@ -217,9 +216,9 @@ func FuzzReadFrame(f *testing.F) {
}

f.Fuzz(func(t *testing.T, input []byte) {
modes := []websocket.Mode{websocket.ClientMode, websocket.ServerMode}
modes := []Mode{ClientMode, ServerMode}
for _, mode := range modes {
frame, err := websocket.ReadFrame(bytes.NewReader(input), mode, 1<<20)
frame, err := ReadFrame(bytes.NewReader(input), mode, 1<<20)
if err != nil {
t.Skipf("skipping eror: %s", err)
return
Expand All @@ -240,9 +239,9 @@ var benchMarkFrameSizes = []int{

func BenchmarkReadFrame(b *testing.B) {
for _, size := range benchMarkFrameSizes {
frame := makeFrame(websocket.OpcodeText, true, size)
frame := makeFrame(OpcodeText, true, size)
buf := &bytes.Buffer{}
assert.NilError(b, websocket.WriteFrame(buf, websocket.NewMaskingKey(), frame))
assert.NilError(b, WriteFrame(buf, NewMaskingKey(), frame))

// Run sub-benchmarks for each payload size
b.Run(formatSize(size), func(b *testing.B) {
Expand All @@ -251,7 +250,7 @@ func BenchmarkReadFrame(b *testing.B) {
b.ResetTimer()
for b.Loop() {
_, _ = src.Seek(0, 0)
frame2, err := websocket.ReadFrame(src, websocket.ServerMode, size)
frame2, err := ReadFrame(src, ServerMode, size)
if err != nil {
b.Fatalf("unexpected error: %v", err)
}
Expand All @@ -264,31 +263,31 @@ func BenchmarkReadFrame(b *testing.B) {
func BenchmarkWriteFrame(b *testing.B) {
for _, size := range benchMarkFrameSizes {
b.Run(formatSize(size), func(b *testing.B) {
frame := makeFrame(websocket.OpcodeText, true, size)
mask := websocket.NewMaskingKey()
frame := makeFrame(OpcodeText, true, size)
mask := NewMaskingKey()
buf := &bytes.Buffer{}

// Write the frame to the buffer once to get the size.
assert.NilError(b, websocket.WriteFrame(buf, mask, frame))
assert.NilError(b, WriteFrame(buf, mask, frame))
expectedSize := len(buf.Bytes())
b.SetBytes(int64(expectedSize))
b.ResetTimer()

for b.Loop() {
buf.Reset()
assert.NilError(b, websocket.WriteFrame(buf, mask, frame))
assert.NilError(b, WriteFrame(buf, mask, frame))
assert.Equal(b, buf.Len(), expectedSize, "payload length")
}
})
}
}

func makeFrame(opcode websocket.Opcode, fin bool, payloadLen int) *websocket.Frame {
func makeFrame(opcode Opcode, fin bool, payloadLen int) *Frame {
payload := make([]byte, payloadLen)
for i := range payload {
payload[i] = 0x20 + byte(i%95) // Map to range 0x20 (space) to 0x7E (~)
}
return websocket.NewFrame(opcode, fin, payload)
return NewFrame(opcode, fin, payload)
}

func formatSize(b int) string {
Expand Down
10 changes: 5 additions & 5 deletions websocket.go
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,7 @@ func HijackConn(w http.ResponseWriter) (net.Conn, error) {
// per the Hijack docs, the returned read buffer may contain unprocessed
// data, so we return a net.Conn implementation that will read from that
// buffer first.
return &BufferedConn{
return &bufferedConn{
Conn: conn,
Reader: rw.Reader,
}, nil
Expand Down Expand Up @@ -538,19 +538,19 @@ func statusCodeForError(err error) (StatusCode, string) {
return StatusInternalError, err.Error()
}

// BufferedConn ties a [net.Conn] to a [bufio.Reader] wrapping that conn, such
// bufferedConn ties a [net.Conn] to a [bufio.Reader] wrapping that conn, such
// as those returned by [http.Hijacker.Hijack], so that all reads go through
// the buffered reader but writes go directly to the underlying conn.
type BufferedConn struct {
type bufferedConn struct {
net.Conn
Reader *bufio.Reader
}

func (bc *BufferedConn) Read(p []byte) (int, error) {
func (bc *bufferedConn) Read(p []byte) (int, error) {
return bc.Reader.Read(p)
}

var _ net.Conn = &BufferedConn{}
var _ net.Conn = &bufferedConn{}

// Hooks define the callbacks that are called during the lifecycle of a
// websocket connection.
Expand Down
Loading
Loading