package proxy import ( "bytes" "testing" ) // maskedFrame builds a masked (client-to-server style) RFC 6455 frame // for opcode carrying payload, mask applied per spec (XOR with a 4-byte // key repeated across the payload). func maskedFrame(t *testing.T, opcode byte, payload []byte, key [4]byte) []byte { t.Helper() var buf bytes.Buffer buf.WriteByte(0x80 | opcode) // FIN=1 if len(payload) > 125 { t.Fatalf("test helper only supports short payloads") } buf.WriteByte(0x80 | byte(len(payload))) // MASK=1 buf.Write(key[:]) masked := make([]byte, len(payload)) for i, b := range payload { masked[i] = b ^ key[i%4] } buf.Write(masked) return buf.Bytes() } // unmaskedFrame builds an unmasked (server-to-client style) frame. func unmaskedFrame(t *testing.T, opcode byte, payload []byte) []byte { t.Helper() var buf bytes.Buffer buf.WriteByte(0x80 | opcode) if len(payload) > 125 { t.Fatalf("test helper only supports short payloads") } buf.WriteByte(byte(len(payload))) // MASK=0 buf.Write(payload) return buf.Bytes() } func TestRelayWSFrameMaskedRoundTrips(t *testing.T) { src := maskedFrame(t, wsOpText, []byte("hello"), [4]byte{1, 2, 3, 4}) var dst bytes.Buffer opcode, payload, err := relayWSFrame(bytes.NewReader(src), &dst) if err != nil { t.Fatalf("unexpected error: %v", err) } if opcode != wsOpText { t.Errorf("opcode = %d, want %d", opcode, wsOpText) } if string(payload) != "hello" { t.Errorf("decoded payload = %q, want %q", payload, "hello") } // The relayed bytes must be byte-for-byte identical to the input - // capture decodes for display, it never re-encodes what's on the // wire. if !bytes.Equal(dst.Bytes(), src) { t.Errorf("relayed bytes = %x, want exactly %x (byte-exact passthrough)", dst.Bytes(), src) } } func TestRelayWSFrameUnmaskedRoundTrips(t *testing.T) { src := unmaskedFrame(t, wsOpBinary, []byte{0xde, 0xad, 0xbe, 0xef}) var dst bytes.Buffer opcode, payload, err := relayWSFrame(bytes.NewReader(src), &dst) if err != nil { t.Fatalf("unexpected error: %v", err) } if opcode != wsOpBinary { t.Errorf("opcode = %d, want %d", opcode, wsOpBinary) } if !bytes.Equal(payload, []byte{0xde, 0xad, 0xbe, 0xef}) { t.Errorf("decoded payload = %x, want deadbeef", payload) } if !bytes.Equal(dst.Bytes(), src) { t.Error("relayed bytes don't match input exactly") } } func TestRelayWSFrameExtended16BitLength(t *testing.T) { payload := bytes.Repeat([]byte("a"), 200) // over 125, forces the 126 extended-length form var buf bytes.Buffer buf.WriteByte(0x80 | wsOpBinary) buf.WriteByte(126) buf.Write([]byte{0x00, 0xc8}) // 200 in 16-bit big-endian buf.Write(payload) var dst bytes.Buffer opcode, got, err := relayWSFrame(bytes.NewReader(buf.Bytes()), &dst) if err != nil { t.Fatalf("unexpected error: %v", err) } if opcode != wsOpBinary { t.Errorf("opcode = %d, want %d", opcode, wsOpBinary) } if !bytes.Equal(got, payload) { t.Errorf("payload length = %d, want %d", len(got), len(payload)) } if !bytes.Equal(dst.Bytes(), buf.Bytes()) { t.Error("relayed bytes don't match input exactly") } } func TestRelayWSFrameOverCapRejected(t *testing.T) { var buf bytes.Buffer buf.WriteByte(0x80 | wsOpBinary) buf.WriteByte(127) lenBytes := make([]byte, 8) big := uint64(maxCaptureBytes + 1) for i := 7; i >= 0; i-- { lenBytes[i] = byte(big) big >>= 8 } buf.Write(lenBytes) // Deliberately no payload bytes written - relayWSFrame must reject // based on the length field alone, before trying to read a body this // large. var dst bytes.Buffer _, _, err := relayWSFrame(bytes.NewReader(buf.Bytes()), &dst) if err == nil { t.Fatal("expected an error for a frame over the capture limit") } } func TestPumpWSCapturesUntilClose(t *testing.T) { var src bytes.Buffer src.Write(unmaskedFrame(t, wsOpText, []byte("one"))) src.Write(unmaskedFrame(t, wsOpText, []byte("two"))) src.Write(unmaskedFrame(t, wsOpClose, nil)) // A frame after close must never be reached. src.Write(unmaskedFrame(t, wsOpText, []byte("unreachable"))) var dst bytes.Buffer var captured [][]byte pumpWS(&src, &dst, func(opcode byte, payload []byte) { captured = append(captured, append([]byte(nil), payload...)) }) if len(captured) != 3 { t.Fatalf("captured %d frames, want 3 (two messages + close)", len(captured)) } if string(captured[0]) != "one" || string(captured[1]) != "two" { t.Errorf("captured = %q, %q, want \"one\", \"two\"", captured[0], captured[1]) } } func TestPumpWSStopsOnReadError(t *testing.T) { // Truncated frame: header claims a payload that never arrives. src := bytes.NewReader([]byte{0x81, 0x05, 'h', 'i'}) var dst bytes.Buffer called := false pumpWS(src, &dst, func(byte, []byte) { called = true }) if called { t.Error("capture should never run for a frame that fails to read fully") } }