From 853da1c9e255a1b05fe0238f56b0898ae75d813e Mon Sep 17 00:00:00 2001 From: Andrey Kolkov Date: Thu, 10 Sep 2026 22:50:37 +0300 Subject: [PATCH] fix: security S1-S4, stress test overflow, SSE slow-client doc MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit S1: ReadTimeout/WriteTimeout (default 0/10s). ReadTimeout opt-in. S2: MaxMessageSize (default 4MB) in readFrame BEFORE allocation. S3: checkSameOrigin by Host only — TLS proxy safe. S4: readFrame uses min(64MB, maxMessageSize). Fix: TestStress_MemoryPressure uint64 overflow on GC. Fix: TestStress_LargeMessages uses MaxMessageSize:16MB. Doc: SSE Hub.Run WARNING about slow-client blocking. All stress tests guarded with testing.Short(). --- sse/hub.go | 5 +++ websocket/conn.go | 64 ++++++++++++++++++++++++++------ websocket/conn_test.go | 24 ++++++------ websocket/export_test.go | 2 +- websocket/frame.go | 19 ++++++---- websocket/frame_test.go | 54 +++++++++++++-------------- websocket/handshake.go | 40 ++++++++++++++------ websocket/handshake_test.go | 12 ++++-- websocket/hub_test.go | 2 +- websocket/load_test.go | 15 ++++++++ websocket/rfc_validation_test.go | 20 +++++----- websocket/stress_test.go | 27 ++++++++++++-- 12 files changed, 196 insertions(+), 88 deletions(-) diff --git a/sse/hub.go b/sse/hub.go index 15b0eb8..6975a56 100644 --- a/sse/hub.go +++ b/sse/hub.go @@ -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]() diff --git a/websocket/conn.go b/websocket/conn.go index 274a960..3295dab 100644 --- a/websocket/conn.go +++ b/websocket/conn.go @@ -6,6 +6,7 @@ import ( "encoding/json/v2" "net" "sync" + "time" "unicode/utf8" ) @@ -54,18 +55,37 @@ type Conn struct { 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, } } @@ -98,8 +118,14 @@ func (c *Conn) Read() (MessageType, []byte, error) { 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 } @@ -147,17 +173,23 @@ func (c *Conn) Read() (MessageType, []byte, error) { 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 { @@ -248,6 +280,12 @@ func (c *Conn) Write(messageType MessageType, data []byte) error { 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 { @@ -266,10 +304,14 @@ func (c *Conn) Write(messageType MessageType, data []byte) error { 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, } diff --git a/websocket/conn_test.go b/websocket/conn_test.go index 3abcc78..0427b19 100644 --- a/websocket/conn_test.go +++ b/websocket/conn_test.go @@ -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). @@ -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. @@ -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 } @@ -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) } @@ -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) } @@ -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) } @@ -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) } @@ -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) } @@ -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) } @@ -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) } @@ -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) } @@ -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)") } diff --git a/websocket/export_test.go b/websocket/export_test.go index f2a7fe9..4c00723 100644 --- a/websocket/export_test.go +++ b/websocket/export_test.go @@ -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 } diff --git a/websocket/frame.go b/websocket/frame.go index 8a1ca51..42cbf2e 100644 --- a/websocket/frame.go +++ b/websocket/frame.go @@ -14,9 +14,8 @@ const ( // 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 @@ -89,7 +88,7 @@ type frame struct { // 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) @@ -154,9 +153,13 @@ func readFrame(r *bufio.Reader) (*frame, error) { 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. @@ -226,7 +229,7 @@ func writeFrame(w *bufio.Writer, f *frame) error { } // 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)) } diff --git a/websocket/frame_test.go b/websocket/frame_test.go index 2ad37c5..692e700 100644 --- a/websocket/frame_test.go +++ b/websocket/frame_test.go @@ -22,7 +22,7 @@ func TestReadFrame_TextUnmasked(t *testing.T) { } r := bufio.NewReader(bytes.NewReader(data)) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("readFrame failed: %v", err) @@ -62,7 +62,7 @@ func TestReadFrame_TextMasked(t *testing.T) { data = append(data, masked...) r := bufio.NewReader(bytes.NewReader(data)) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("readFrame failed: %v", err) @@ -91,7 +91,7 @@ func TestReadFrame_Binary(t *testing.T) { data = append(data, payload...) r := bufio.NewReader(bytes.NewReader(data)) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("readFrame failed: %v", err) @@ -149,7 +149,7 @@ func TestReadFrame_Fragmented(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { r := bufio.NewReader(bytes.NewReader(tt.data)) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("readFrame failed: %v", err) @@ -201,7 +201,7 @@ func TestReadFrame_ControlFrames(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { r := bufio.NewReader(bytes.NewReader(tt.data)) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("readFrame failed: %v", err) @@ -235,7 +235,7 @@ func TestReadFrame_ExtendedLength16(t *testing.T) { data = append(data, payload...) r := bufio.NewReader(bytes.NewReader(data)) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("readFrame failed: %v", err) @@ -264,7 +264,7 @@ func TestReadFrame_ExtendedLength64(t *testing.T) { data = append(data, payload...) r := bufio.NewReader(bytes.NewReader(data)) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("readFrame failed: %v", err) @@ -288,7 +288,7 @@ func TestReadFrame_InvalidOpcode(t *testing.T) { } r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if !errors.Is(err, ErrInvalidOpcode) { t.Errorf("expected ErrInvalidOpcode, got %v", err) @@ -314,7 +314,7 @@ func TestReadFrame_ReservedBits(t *testing.T) { data := []byte{tt.byte0, 0x00} r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if !errors.Is(err, ErrReservedBits) { t.Errorf("expected ErrReservedBits, got %v", err) @@ -333,7 +333,7 @@ func TestReadFrame_ControlFragmented(t *testing.T) { } r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if !errors.Is(err, ErrControlFragmented) { t.Errorf("expected ErrControlFragmented, got %v", err) @@ -352,7 +352,7 @@ func TestReadFrame_ControlTooLarge(t *testing.T) { data = append(data, make([]byte, 126)...) r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if !errors.Is(err, ErrControlTooLarge) { t.Errorf("expected ErrControlTooLarge, got %v", err) @@ -372,7 +372,7 @@ func TestReadFrame_InvalidUTF8(t *testing.T) { data = append(data, invalidUTF8...) r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if !errors.Is(err, ErrInvalidUTF8) { t.Errorf("expected ErrInvalidUTF8, got %v", err) @@ -690,7 +690,7 @@ func TestRoundTrip(t *testing.T) { // Read frame. r := bufio.NewReader(&buf) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("readFrame failed: %v", err) @@ -785,7 +785,7 @@ func TestReadFrame_IncompleteHeader(t *testing.T) { data := []byte{0x81} r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if err == nil { t.Error("expected error for incomplete header") @@ -805,7 +805,7 @@ func TestReadFrame_IncompletePayload(t *testing.T) { } r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if err == nil { t.Error("expected error for incomplete payload") @@ -900,7 +900,7 @@ func BenchmarkReadFrame_Small(b *testing.B) { for i := 0; i < b.N; i++ { r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if err != nil { b.Fatal(err) } @@ -923,7 +923,7 @@ func BenchmarkReadFrame_Medium(b *testing.B) { for i := 0; i < b.N; i++ { r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if err != nil { b.Fatal(err) } @@ -946,7 +946,7 @@ func BenchmarkReadFrame_Large(b *testing.B) { for i := 0; i < b.N; i++ { r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if err != nil { b.Fatal(err) } @@ -1093,7 +1093,7 @@ func TestUTF8Validation(t *testing.T) { data = append(data, tt.payload...) r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if tt.valid && err != nil { t.Errorf("expected no error for valid UTF-8, got %v", err) @@ -1108,7 +1108,7 @@ func TestUTF8Validation(t *testing.T) { // TestMaxPayloadLength tests maximum payload length enforcement. func TestMaxPayloadLength(t *testing.T) { // Test data frame at limit. - payloadLen := maxFramePayload + payloadLen := defaultMaxFramePayload data := []byte{0x82, 127} // Binary, 64-bit length lenBuf := make([]byte, 8) binary.BigEndian.PutUint64(lenBuf, uint64(payloadLen)) @@ -1116,7 +1116,7 @@ func TestMaxPayloadLength(t *testing.T) { // Don't actually create huge payload, just test header. r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) // Should fail on reading payload (EOF), but not on size validation. if err == nil { @@ -1139,7 +1139,7 @@ func TestFragmentationSequence(t *testing.T) { for i, frameData := range frames { r := bufio.NewReader(bytes.NewReader(frameData)) - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("frame %d: readFrame failed: %v", i, err) @@ -1171,7 +1171,7 @@ func TestReadFrame_MSBSet(t *testing.T) { } r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if !errors.Is(err, ErrProtocolError) { t.Errorf("expected ErrProtocolError for MSB=1, got %v", err) @@ -1213,7 +1213,7 @@ func TestReadFrame_IncompleteMask(t *testing.T) { } r := bufio.NewReader(bytes.NewReader(data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if err == nil { t.Error("expected error for incomplete mask") @@ -1250,7 +1250,7 @@ func TestReadFrame_IncompleteExtendedLength(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { r := bufio.NewReader(bytes.NewReader(tt.data)) - _, err := readFrame(r) + _, err := readFrame(r, 0) if err == nil { t.Error("expected error for incomplete extended length") @@ -1264,8 +1264,8 @@ func TestReadFrame_IncompleteExtendedLength(t *testing.T) { // TestWriteFrame_FrameTooLarge tests payload exceeding max size. func TestWriteFrame_FrameTooLarge(t *testing.T) { - // Create payload exceeding maxFramePayload. - payloadLen := maxFramePayload + 1 + // Create payload exceeding defaultMaxFramePayload. + payloadLen := defaultMaxFramePayload + 1 f := &frame{ fin: true, diff --git a/websocket/handshake.go b/websocket/handshake.go index 907ab52..4419e8d 100644 --- a/websocket/handshake.go +++ b/websocket/handshake.go @@ -5,7 +5,9 @@ import ( "crypto/sha1" // #nosec G505 - SHA-1 required by RFC 6455 Section 1.3 "encoding/base64" "net/http" + "net/url" "strings" + "time" ) // Magic GUID from RFC 6455 Section 1.3. @@ -43,8 +45,20 @@ type UpgradeOptions struct { ReadBufferSize int // WriteBufferSize sets size of write buffer (default: 4096). - // Larger buffers reduce syscalls for large messages. WriteBufferSize int + + // MaxMessageSize limits total reassembled message size (default: 4 MB). + // Prevents memory exhaustion from fragmented messages. + // Set to 0 for no limit (not recommended). + MaxMessageSize int64 + + // ReadTimeout sets deadline for read operations (default: 60s). + // Zero means no timeout (not recommended in production). + ReadTimeout time.Duration + + // WriteTimeout sets deadline for write operations (default: 10s). + // Zero means no timeout (not recommended in production). + WriteTimeout time.Duration } // Upgrade upgrades an HTTP connection to the WebSocket protocol. @@ -121,8 +135,13 @@ func Upgrade(w http.ResponseWriter, r *http.Request, opts *UpgradeOptions) (*Con return nil, ErrMissingSecKey } - // 6. Check origin (application-level security) - if opts.CheckOrigin != nil && !opts.CheckOrigin(r) { + // 6. Check origin (application-level security). + // Default: same-origin check. Set CheckOrigin to override. + checkOrigin := opts.CheckOrigin + if checkOrigin == nil { + checkOrigin = checkSameOrigin + } + if !checkOrigin(r) { return nil, ErrOriginDenied } @@ -171,7 +190,7 @@ func Upgrade(w http.ResponseWriter, r *http.Request, opts *UpgradeOptions) (*Con writer := bufio.NewWriterSize(netConn, opts.WriteBufferSize) // 12. Create WebSocket connection (server-side) - conn := newConn(netConn, reader, writer, true) + conn := newConn(netConn, reader, writer, true, *opts) return conn, nil } @@ -257,13 +276,12 @@ func checkSameOrigin(r *http.Request) bool { return true } - // Build expected origin from request - scheme := "http" - if r.TLS != nil { - scheme = "https" + // Compare host only, ignore scheme — works behind TLS-terminating proxies + // (nginx/Caddy/Cloudflare set r.TLS=nil even for HTTPS traffic). + u, err := url.Parse(origin) + if err != nil { + return false } - expectedOrigin := scheme + "://" + r.Host - - return origin == expectedOrigin + return strings.EqualFold(u.Host, r.Host) } diff --git a/websocket/handshake_test.go b/websocket/handshake_test.go index 8015f9e..17ac58d 100644 --- a/websocket/handshake_test.go +++ b/websocket/handshake_test.go @@ -206,10 +206,16 @@ func TestUpgrade_OriginCheck(t *testing.T) { wantErr error }{ { - name: "no check - allow all", + name: "nil CheckOrigin defaults to same-origin - rejects cross-origin", origin: "http://evil.com", checkOrigin: nil, - wantErr: ErrHijackFailed, // Will fail at hijack + wantErr: ErrOriginDenied, + }, + { + name: "explicit allow-all accepts any origin", + origin: "http://evil.com", + checkOrigin: func(_ *http.Request) bool { return true }, + wantErr: ErrHijackFailed, // passes origin → fails at hijack }, { name: "check passes", @@ -593,7 +599,7 @@ func TestCheckSameOrigin(t *testing.T) { origin: "https://example.com", host: "example.com", tls: false, - want: false, + want: true, // scheme ignored for TLS-terminating proxy compatibility }, } diff --git a/websocket/hub_test.go b/websocket/hub_test.go index 9ac3308..2f9ef78 100644 --- a/websocket/hub_test.go +++ b/websocket/hub_test.go @@ -421,7 +421,7 @@ func (c *mockHubClient) extractMessages() { // Read frame from buffer reader := bufio.NewReader(bytes.NewReader(c.writeBuf.Bytes())) - frame, err := readFrame(reader) + frame, err := readFrame(reader, 0) if err != nil { c.mu.Unlock() continue diff --git a/websocket/load_test.go b/websocket/load_test.go index b72ae58..f0690f9 100644 --- a/websocket/load_test.go +++ b/websocket/load_test.go @@ -14,6 +14,9 @@ import ( // TestLoad_ConcurrentConnections tests handling 100 concurrent WebSocket connections. func TestLoad_ConcurrentConnections(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } if testing.Short() { t.Skip("Skipping load test in short mode") } @@ -135,6 +138,9 @@ func TestLoad_ConcurrentConnections(t *testing.T) { // TestLoad_Hub_100Clients tests Hub broadcasting to 100 concurrent clients. func TestLoad_Hub_100Clients(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } if testing.Short() { t.Skip("Skipping load test in short mode") } @@ -269,6 +275,9 @@ func TestLoad_Hub_100Clients(t *testing.T) { // TestLoad_SSE_100Clients tests SSE Broker broadcasting to 100 concurrent clients. // This test is placed in websocket package to compare performance with WebSocket Hub. func TestLoad_SSE_100Clients(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } if testing.Short() { t.Skip("Skipping load test in short mode") } @@ -281,6 +290,9 @@ func TestLoad_SSE_100Clients(t *testing.T) { // TestLoad_RapidMessages tests rapid message sending and receiving. func TestLoad_RapidMessages(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } if testing.Short() { t.Skip("Skipping load test in short mode") } @@ -382,6 +394,9 @@ func TestLoad_RapidMessages(t *testing.T) { // TestLoad_ParallelHubs tests multiple Hubs running concurrently. func TestLoad_ParallelHubs(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } if testing.Short() { t.Skip("Skipping load test in short mode") } diff --git a/websocket/rfc_validation_test.go b/websocket/rfc_validation_test.go index 343cca9..67830d3 100644 --- a/websocket/rfc_validation_test.go +++ b/websocket/rfc_validation_test.go @@ -64,7 +64,7 @@ func TestRFC_ControlFramesDuringFragmentation(t *testing.T) { r := bufio.NewReader(&buf) // Read first fragment - frame1, err := readFrame(r) + frame1, err := readFrame(r, 0) if err != nil { t.Fatalf("Read fragment 1 failed: %v", err) } @@ -76,7 +76,7 @@ func TestRFC_ControlFramesDuringFragmentation(t *testing.T) { } // Read PING (control frame during fragmentation) - pingFrame, err := readFrame(r) + pingFrame, err := readFrame(r, 0) if err != nil { t.Fatalf("Read PING failed: %v", err) } @@ -88,7 +88,7 @@ func TestRFC_ControlFramesDuringFragmentation(t *testing.T) { } // Read continuation frame - frame2, err := readFrame(r) + frame2, err := readFrame(r, 0) if err != nil { t.Fatalf("Read continuation failed: %v", err) } @@ -97,7 +97,7 @@ func TestRFC_ControlFramesDuringFragmentation(t *testing.T) { } // Read final continuation - frame3, err := readFrame(r) + frame3, err := readFrame(r, 0) if err != nil { t.Fatalf("Read final continuation failed: %v", err) } @@ -151,7 +151,7 @@ func TestRFC_PayloadLengthBoundaries(t *testing.T) { // Read frame back r := bufio.NewReader(&buf) - readBack, err := readFrame(r) + readBack, err := readFrame(r, 0) if err != nil { t.Fatalf("Read failed: %v", err) } @@ -304,7 +304,7 @@ func TestRFC_UTF8Validation_Extended(t *testing.T) { // Try to read back r := bufio.NewReader(&buf) - _, err := readFrame(r) + _, err := readFrame(r, 0) if tt.wantError && err == nil { t.Error("Expected read to fail with invalid UTF-8, but it succeeded") } @@ -344,7 +344,7 @@ func TestRFC_FragmentationSequence(t *testing.T) { r := bufio.NewReader(&buf) // First frame: FIN=0, opcode=text - f1, err := readFrame(r) + f1, err := readFrame(r, 0) if err != nil { t.Fatalf("Read first frame failed: %v", err) } @@ -354,7 +354,7 @@ func TestRFC_FragmentationSequence(t *testing.T) { // Continuation frames: FIN=0, opcode=continuation for i := 1; i < 3; i++ { - f, err := readFrame(r) + f, err := readFrame(r, 0) if err != nil { t.Fatalf("Read continuation %d failed: %v", i, err) } @@ -364,7 +364,7 @@ func TestRFC_FragmentationSequence(t *testing.T) { } // Final frame: FIN=1, opcode=continuation - fFinal, err := readFrame(r) + fFinal, err := readFrame(r, 0) if err != nil { t.Fatalf("Read final frame failed: %v", err) } @@ -417,7 +417,7 @@ func TestRFC_CloseFramePayload(t *testing.T) { // Read back r := bufio.NewReader(&buf) - readBack, err := readFrame(r) + readBack, err := readFrame(r, 0) if err != nil { t.Fatalf("Read failed: %v", err) } diff --git a/websocket/stress_test.go b/websocket/stress_test.go index ececa92..2f67d8e 100644 --- a/websocket/stress_test.go +++ b/websocket/stress_test.go @@ -17,12 +17,13 @@ import ( // TestStress_LargeMessages tests handling of large messages (fragmented). func TestStress_LargeMessages(t *testing.T) { if testing.Short() { - t.Skip("Skipping stress test in short mode") + t.Skip("skipping stress test in short mode") } - // Setup echo server + largeOpts := &UpgradeOptions{MaxMessageSize: 16 * 1024 * 1024} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - conn, err := Upgrade(w, r, nil) + conn, err := Upgrade(w, r, largeOpts) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return @@ -100,6 +101,9 @@ func TestStress_LargeMessages(t *testing.T) { // TestStress_RapidConnectDisconnect tests rapid connection cycling. func TestStress_RapidConnectDisconnect(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } if testing.Short() { t.Skip("Skipping stress test in short mode") } @@ -225,6 +229,9 @@ func TestStress_RapidConnectDisconnect(t *testing.T) { // TestStress_ConcurrentBroadcast tests concurrent broadcasting from multiple goroutines. func TestStress_ConcurrentBroadcast(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } if testing.Short() { t.Skip("Skipping stress test in short mode") } @@ -384,6 +391,9 @@ func TestStress_ConcurrentBroadcast(t *testing.T) { // TestStress_MemoryPressure tests behavior under memory pressure with many concurrent operations. func TestStress_MemoryPressure(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } if testing.Short() { t.Skip("Skipping stress test in short mode") } @@ -471,7 +481,10 @@ func TestStress_MemoryPressure(t *testing.T) { runtime.ReadMemStats(&memStatsAfter) // Memory metrics - allocIncrease := memStatsAfter.Alloc - memStatsBefore.Alloc + var allocIncrease uint64 + if memStatsAfter.Alloc > memStatsBefore.Alloc { + allocIncrease = memStatsAfter.Alloc - memStatsBefore.Alloc + } totalAllocIncrease := memStatsAfter.TotalAlloc - memStatsBefore.TotalAlloc t.Logf("Memory metrics:") @@ -491,11 +504,17 @@ func TestStress_MemoryPressure(t *testing.T) { // TestStress_PingPongStorm tests handling of many ping/pong control frames. // NOTE: Skipped - requires SetPongHandler() and WritePing() methods not yet implemented. func TestStress_PingPongStorm(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } t.Skip("Requires SetPongHandler() and WritePing() methods - TODO") } // TestStress_ConnectionTimeout tests handling of connection timeouts and deadlines. // NOTE: Skipped - requires SetReadDeadline() method not yet implemented. func TestStress_ConnectionTimeout(t *testing.T) { + if testing.Short() { + t.Skip("skipping stress test in short mode") + } t.Skip("Requires SetReadDeadline() method - TODO") }