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
5 changes: 5 additions & 0 deletions sse/hub.go
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,11 @@ func NewHub[T any]() *Hub[T] {
// Run processes client registration, unregistration, and broadcast operations.
// It should be called in a goroutine and will block until Close() is called.
//
// WARNING: Broadcast writes to clients synchronously in the Run loop.
// A single slow or disconnected client will block delivery to all others.
// For production use, set WriteTimeout on your http.Server or use
// http.ResponseController.SetWriteDeadline (Go 1.20+) per connection.
//
// Example:
//
// hub := sse.NewHub[string]()
Expand Down
64 changes: 53 additions & 11 deletions websocket/conn.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
"encoding/json/v2"
"net"
"sync"
"time"
"unicode/utf8"
)

Expand Down Expand Up @@ -54,18 +55,37 @@
fragmentBuf bytes.Buffer // Accumulates fragmented message
fragmentType byte // Opcode of first fragment (text/binary)
inFragment bool // Currently reading fragmented message

// Limits and timeouts
maxMessageSize int64 // 0 = no limit
readTimeout time.Duration // 0 = no timeout
writeTimeout time.Duration // 0 = no timeout
}

// newConn creates a new WebSocket connection (internal constructor).
//
// Called by Upgrade() after successful handshake.
// Not exported - users should call Upgrade() to create connections.
func newConn(netConn net.Conn, reader *bufio.Reader, writer *bufio.Writer, isServer bool) *Conn {
func newConn(netConn net.Conn, reader *bufio.Reader, writer *bufio.Writer, isServer bool, opts UpgradeOptions) *Conn {
maxMsg := opts.MaxMessageSize
if maxMsg == 0 {
maxMsg = 4 * 1024 * 1024 // 4 MB default
}
readTO := opts.ReadTimeout
// No default read timeout — server has no built-in ping loop.
// Set ReadTimeout explicitly alongside your own PingPeriod.
writeTO := opts.WriteTimeout
if writeTO == 0 {
writeTO = 10 * time.Second
}
return &Conn{
conn: netConn,
reader: reader,
writer: writer,
isServer: isServer,
conn: netConn,
reader: reader,
writer: writer,
isServer: isServer,
maxMessageSize: maxMsg,
readTimeout: readTO,
writeTimeout: writeTO,
}
}

Expand Down Expand Up @@ -98,8 +118,14 @@
c.closeMu.RUnlock()

for {
if c.readTimeout > 0 && c.conn != nil {
if err := c.conn.SetReadDeadline(time.Now().Add(c.readTimeout)); err != nil {
return 0, nil, err
}
}

// Read next frame
f, err := readFrame(c.reader)
f, err := readFrame(c.reader, c.maxMessageSize)
if err != nil {
return 0, nil, err
}
Expand Down Expand Up @@ -147,17 +173,23 @@
c.inFragment = true
c.fragmentType = f.opcode
c.fragmentBuf.Reset()

if c.maxMessageSize > 0 && int64(len(f.payload)) > c.maxMessageSize {
_ = c.CloseWithCode(CloseMessageTooBig, "message too big")
return 0, nil, ErrMessageTooLarge
}
c.fragmentBuf.Write(f.payload)

case opcodeContinuation:
// Continuation frame
if !c.inFragment {
// Unexpected continuation (no prior fragment)
_ = c.CloseWithCode(CloseProtocolError, "unexpected continuation")
return 0, nil, ErrUnexpectedContinuation
}

// Append to fragment buffer
if c.maxMessageSize > 0 && int64(c.fragmentBuf.Len())+int64(len(f.payload)) > c.maxMessageSize {
_ = c.CloseWithCode(CloseMessageTooBig, "message too big")
return 0, nil, ErrMessageTooLarge
}
c.fragmentBuf.Write(f.payload)

if f.fin {
Expand Down Expand Up @@ -248,6 +280,12 @@
c.writeMu.Lock()
defer c.writeMu.Unlock()

if c.writeTimeout > 0 && c.conn != nil {
if err := c.conn.SetWriteDeadline(time.Now().Add(c.writeTimeout)); err != nil {
return err
}
}

// Build frame
var opcode byte
switch messageType {
Expand All @@ -266,10 +304,14 @@
return ErrInvalidMessageType
}

if c.maxMessageSize > 0 && int64(len(data)) > c.maxMessageSize {
return ErrMessageTooLarge
}

f := &frame{
fin: true, // Single frame (no fragmentation yet)
fin: true,
opcode: opcode,
masked: !c.isServer, // Server: NO mask, Client: YES mask
masked: !c.isServer,
payload: data,
}

Expand Down Expand Up @@ -420,7 +462,7 @@

// Build close frame payload: 2 bytes status code + optional reason
payload := make([]byte, 2+len(reason))
payload[0] = byte(code >> 8)

Check failure on line 465 in websocket/conn.go

View workflow job for this annotation

GitHub Actions / Lint

G115: integer overflow conversion int -> byte (gosec)
payload[1] = byte(code & 0xFF)
copy(payload[2:], reason)

Expand Down
24 changes: 12 additions & 12 deletions websocket/conn_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ func mockConn(t *testing.T, frames []*frame, isServer bool) *Conn {
// Create connection with buffer as reader
reader := bufio.NewReader(&buf)
writer := bufio.NewWriter(io.Discard) // Writes go nowhere
return newConn(nil, reader, writer, isServer)
return newConn(nil, reader, writer, isServer, UpgradeOptions{})
}

// mockConnNoValidation creates a mock connection with frames (no validation).
Expand All @@ -50,7 +50,7 @@ func mockConnNoValidation(t *testing.T, frames []*frame, isServer bool) *Conn {
// Create connection with buffer as reader
reader := bufio.NewReader(&buf)
writer := bufio.NewWriter(io.Discard) // Writes go nowhere
return newConn(nil, reader, writer, isServer)
return newConn(nil, reader, writer, isServer, UpgradeOptions{})
}

// mockConnWriter creates a mock connection that captures writes.
Expand All @@ -62,7 +62,7 @@ func mockConnWriter(t *testing.T) (*Conn, *bytes.Buffer) {
var writeBuf bytes.Buffer
reader := bufio.NewReader(bytes.NewReader(nil)) // Empty reader
writer := bufio.NewWriter(&writeBuf)
conn := newConn(nil, reader, writer, true) // Server-side
conn := newConn(nil, reader, writer, true, UpgradeOptions{}) // Server-side
return conn, &writeBuf
}

Expand Down Expand Up @@ -346,7 +346,7 @@ func TestConn_Write(t *testing.T) {

// Read frame from buffer
r := bufio.NewReader(writeBuf)
frame, err := readFrame(r)
frame, err := readFrame(r, 0)
if err != nil {
t.Fatalf("readFrame() error = %v", err)
}
Expand Down Expand Up @@ -377,7 +377,7 @@ func TestConn_WriteText(t *testing.T) {
}

r := bufio.NewReader(writeBuf)
frame, err := readFrame(r)
frame, err := readFrame(r, 0)
if err != nil {
t.Fatalf("readFrame() error = %v", err)
}
Expand Down Expand Up @@ -407,7 +407,7 @@ func TestConn_WriteJSON(t *testing.T) {
}

r := bufio.NewReader(writeBuf)
frame, err := readFrame(r)
frame, err := readFrame(r, 0)
if err != nil {
t.Fatalf("readFrame() error = %v", err)
}
Expand Down Expand Up @@ -437,7 +437,7 @@ func TestConn_Ping(t *testing.T) {
}

r := bufio.NewReader(writeBuf)
frame, err := readFrame(r)
frame, err := readFrame(r, 0)
if err != nil {
t.Fatalf("readFrame() error = %v", err)
}
Expand Down Expand Up @@ -466,7 +466,7 @@ func TestConn_Pong(t *testing.T) {
}

r := bufio.NewReader(writeBuf)
frame, err := readFrame(r)
frame, err := readFrame(r, 0)
if err != nil {
t.Fatalf("readFrame() error = %v", err)
}
Expand Down Expand Up @@ -495,7 +495,7 @@ func TestConn_Close(t *testing.T) {

// Verify close frame sent
r := bufio.NewReader(writeBuf)
frame, err := readFrame(r)
frame, err := readFrame(r, 0)
if err != nil {
t.Fatalf("readFrame() error = %v", err)
}
Expand Down Expand Up @@ -538,7 +538,7 @@ func TestConn_CloseWithCode(t *testing.T) {

// Verify close frame
r := bufio.NewReader(writeBuf)
frame, err := readFrame(r)
frame, err := readFrame(r, 0)
if err != nil {
t.Fatalf("readFrame() error = %v", err)
}
Expand Down Expand Up @@ -609,7 +609,7 @@ func TestConn_DoubleClose(t *testing.T) {

// Read first close frame
r := bufio.NewReader(writeBuf)
frame1, err := readFrame(r)
frame1, err := readFrame(r, 0)
if err != nil {
t.Fatalf("readFrame() error = %v", err)
}
Expand All @@ -624,7 +624,7 @@ func TestConn_DoubleClose(t *testing.T) {
}

// Try to read second frame (should be EOF)
frame2, err := readFrame(r)
frame2, err := readFrame(r, 0)
if err == nil && frame2 != nil {
t.Error("Second close frame sent (Close not idempotent)")
}
Expand Down
2 changes: 1 addition & 1 deletion websocket/export_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ type FrameForTest struct {

// ReadFrameForTest reads a frame (exported for testing).
func ReadFrameForTest(r *bufio.Reader) (*FrameForTest, error) {
f, err := readFrame(r)
f, err := readFrame(r, 0)
if err != nil {
return nil, err
}
Expand Down
19 changes: 11 additions & 8 deletions websocket/frame.go
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,8 @@
// RFC 6455 Section 5.5: Control frames must have payload <= 125 bytes.
maxControlPayload = 125

// maxFramePayload is the maximum payload length for data frames.
// Default: 32 MB (configurable in production).
maxFramePayload = 32 * 1024 * 1024
// defaultMaxFramePayload is the default maximum payload for a single data frame.
defaultMaxFramePayload = 64 * 1024 * 1024

// Payload length encoding thresholds (RFC 6455 Section 5.2).
payloadLen7Bit = 125 // 0-125: stored in 7 bits
Expand Down Expand Up @@ -89,7 +88,7 @@
// Returns:
// - frame: parsed frame structure
// - error: validation or I/O error
func readFrame(r *bufio.Reader) (*frame, error) {
func readFrame(r *bufio.Reader, maxMessageSize int64) (*frame, error) {
// Step 1: Read 2-byte header.
// Byte 0: FIN(1) RSV(3) Opcode(4)
// Byte 1: MASK(1) PayloadLen(7)
Expand Down Expand Up @@ -154,9 +153,13 @@
return nil, ErrControlTooLarge
}

// Validate data frame payload length (implementation limit).
if payloadLen > maxFramePayload {
return nil, fmt.Errorf("%w: %d bytes", ErrFrameTooLarge, payloadLen)
// Validate data frame payload length BEFORE allocating.
limit := uint64(defaultMaxFramePayload)
if maxMessageSize > 0 && uint64(maxMessageSize) < limit {
limit = uint64(maxMessageSize)
}
if payloadLen > limit {
return nil, fmt.Errorf("%w: %d bytes (limit %d)", ErrFrameTooLarge, payloadLen, limit)
}

// Step 3: Read masking key if MASK=1.
Expand Down Expand Up @@ -226,7 +229,7 @@
}

// Validate payload length (implementation limit).
if len(f.payload) > maxFramePayload {
if len(f.payload) > defaultMaxFramePayload {
return fmt.Errorf("%w: %d bytes", ErrFrameTooLarge, len(f.payload))
}

Expand Down Expand Up @@ -256,7 +259,7 @@
payloadLen := uint64(len(f.payload))

// Determine payload length encoding.
//nolint:gosec // G602: False positive - header is always length 2

Check failure on line 262 in websocket/frame.go

View workflow job for this annotation

GitHub Actions / Lint

directive `//nolint:gosec // G602: False positive - header is always length 2` is unused for linter "gosec" (nolintlint)
switch {
case payloadLen <= payloadLen7Bit:
// 7-bit length (0-125).
Expand Down Expand Up @@ -360,7 +363,7 @@
payloadLen := uint64(len(f.payload))

// Determine payload length encoding.
//nolint:gosec // G602: False positive - header is always length 2

Check failure on line 366 in websocket/frame.go

View workflow job for this annotation

GitHub Actions / Lint

directive `//nolint:gosec // G602: False positive - header is always length 2` is unused for linter "gosec" (nolintlint)
switch {
case payloadLen <= payloadLen7Bit:
// 7-bit length (0-125).
Expand Down
Loading
Loading