srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/internal/proxy/websocket_test.go
blob: 8256567f3370b4b63c5e455500027f8ef6ec6f8f (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
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")
	}
}