diff options
Diffstat (limited to 'internal/proxy/websocket_test.go')
| -rw-r--r-- | internal/proxy/websocket_test.go | 159 |
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") + } +} |