srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/internal/proxy/websocket_test.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/proxy/websocket_test.go')
-rw-r--r--internal/proxy/websocket_test.go159
1 files changed, 159 insertions, 0 deletions
diff --git a/internal/proxy/websocket_test.go b/internal/proxy/websocket_test.go
new file mode 100644
index 0000000..8256567
--- /dev/null
+++ b/internal/proxy/websocket_test.go
@@ -0,0 +1,159 @@
+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")
+ }
+}