ws_frames_test.mx raw
1 package ws
2
3 import (
4 "syscall"
5 "testing"
6 )
7
8 // Framing and handshake paths the subscribe test does not reach: both extended
9 // length encodings, a masked server frame, ping/pong/close control frames,
10 // protocol and truncation errors, the client's own masked write, and the
11 // handshake rejections.
12 //
13 // Each case runs a loopback server in a spawn domain that writes a scripted
14 // byte sequence after the upgrade, so no real relay and no TLS are involved.
15
16 // serveScript completes the handshake and writes frames verbatim.
17 func serveScript(fd int32, frames []byte) {
18 nfd, _, err := syscall.Accept(fd)
19 if err != nil {
20 syscall.Close(fd)
21 return
22 }
23 req := readUpgradeRequest(nfd)
24 syscall.Write(nfd, []byte(upgradeHead(req)))
25 if len(frames) > 0 {
26 syscall.Write(nfd, frames)
27 }
28 // A short pause keeps the bytes in flight before the FIN; the client reads
29 // what is buffered either way.
30 syscall.Close(nfd)
31 syscall.Close(fd)
32 }
33
34 // serveReply writes a raw HTTP reply to the upgrade request (for Dial errors).
35 func serveReply(fd int32, reply []byte) {
36 nfd, _, err := syscall.Accept(fd)
37 if err != nil {
38 syscall.Close(fd)
39 return
40 }
41 readUpgradeRequest(nfd)
42 syscall.Write(nfd, reply)
43 syscall.Close(nfd)
44 syscall.Close(fd)
45 }
46
47 // serveEchoFrame reads count masked client frames and echoes each payload as
48 // an unmasked text frame, which proves the client masked it correctly.
49 func serveEchoFrame(fd int32, count int32) {
50 nfd, _, err := syscall.Accept(fd)
51 if err != nil {
52 syscall.Close(fd)
53 return
54 }
55 req := readUpgradeRequest(nfd)
56 syscall.Write(nfd, []byte(upgradeHead(req)))
57 for i := int32(0); i < count; i++ {
58 body := readClientFrame(nfd)
59 syscall.Write(nfd, serverFrame(OpText, body))
60 }
61 syscall.Close(nfd)
62 syscall.Close(fd)
63 }
64
65 // readN reads exactly n bytes or returns what it got.
66 func readN(fd int32, n int32) (out []byte) {
67 out = []byte{:0:n}
68 for len(out) < n {
69 buf := []byte{:n - len(out)}
70 got, err := syscall.Read(fd, buf)
71 if got <= 0 || err != nil {
72 return out
73 }
74 out = out | buf[:got]
75 }
76 return
77 }
78
79 // readClientFrame decodes one client frame (always masked) and unmasks it.
80 func readClientFrame(fd int32) (payload []byte) {
81 hdr := readN(fd, 2)
82 if len(hdr) < 2 {
83 return nil
84 }
85 plen := int32(hdr[1] & 0x7F)
86 if plen == 126 {
87 ext := readN(fd, 2)
88 if len(ext) < 2 {
89 return nil
90 }
91 plen = int32(ext[0])<<8 | int32(ext[1])
92 } else if plen == 127 {
93 ext := readN(fd, 8)
94 if len(ext) < 8 {
95 return nil
96 }
97 plen = int32(ext[4])<<24 | int32(ext[5])<<16 | int32(ext[6])<<8 | int32(ext[7])
98 }
99 var mask [4]byte
100 if hdr[1]&0x80 != 0 {
101 m := readN(fd, 4)
102 if len(m) < 4 {
103 return nil
104 }
105 mask = [4]byte{m[0], m[1], m[2], m[3]}
106 }
107 body := readN(fd, plen)
108 for i := int32(0); i < int32(len(body)); i++ {
109 body[i] = body[i] ^ mask[i%4]
110 }
111 return body
112 }
113
114 // rawFrame builds an unmasked frame with an explicit length byte, so a test can
115 // exercise the 126 and 127 encodings and the masked variant.
116 func rawFrame(op byte, plen int32, extended []byte, payload []byte) (buf []byte) {
117 buf = []byte{:2}
118 buf[0] = 0x80 | op
119 buf[1] = byte(plen)
120 buf = buf | extended
121 return buf | payload
122 }
123
124 func dialScript(t *testing.T, frames []byte) (c *Conn, done chan struct{}) {
125 t.Helper()
126 fd, port, err := listenLoopback()
127 if err != nil {
128 t.Fatalf("listen: %v", err)
129 return
130 }
131 done = spawn(serveScript, fd, frames)
132 c, cerr := Dial("ws://127.0.0.1:" | portStr(port) | "/")
133 if cerr != nil {
134 t.Fatalf("dial: %v", cerr)
135 return
136 }
137 return
138 }
139
140 func portStr(port int32) (s string) {
141 if port == 0 {
142 return "0"
143 }
144 buf := []byte{:6}
145 n := int32(6)
146 for port > 0 {
147 n--
148 buf[n] = byte('0' + port%10)
149 port = port / 10
150 }
151 return string(buf[n:])
152 }
153
154 func TestReadMessageLengthEncodings(t *testing.T) {
155 // 126: 200 bytes, extended 16-bit length.
156 small := makePayload(200)
157 cases := []struct {
158 name string
159 op byte
160 frame []byte
161 want []byte
162 }{
163 {"short", OpText, rawFrame(OpText, 5, nil, []byte("hello")), []byte("hello")},
164 {"extended16", OpText,
165 rawFrame(OpText, 126, []byte{byte(200 >> 8), byte(200)}, small), small},
166 {"extended64", OpText,
167 rawFrame(OpText, 127, []byte{0, 0, 0, 0,
168 byte(70000 >> 24), byte(70000 >> 16), byte(70000 >> 8), byte(70000)}, makePayload(70000)),
169 makePayload(70000)},
170 {"binary", OpBinary, rawFrame(OpBinary, 3, nil, []byte{1, 2, 3}), []byte{1, 2, 3}},
171 {"empty", OpText, rawFrame(OpText, 0, nil, nil), nil},
172 }
173 for _, tc := range cases {
174 c, done := dialScript(t, tc.frame)
175 if c == nil {
176 continue
177 }
178 op, payload, err := c.ReadMessage()
179 if err != nil {
180 t.Fatalf("%s: %s", tc.name, err.Error())
181 }
182 if op != tc.op {
183 t.Fatalf("%s: op = %d, want %d", tc.name, int32(op), int32(tc.op))
184 }
185 if string(payload) != string(tc.want) {
186 t.Fatalf("%s: payload length %d, want %d", tc.name, len(payload), len(tc.want))
187 }
188 c.Close()
189 <-done
190 }
191 }
192
193 func makePayload(n int32) (b []byte) {
194 b = []byte{:n}
195 for i := int32(0); i < n; i++ {
196 b[i] = byte('a' + i%26)
197 }
198 return
199 }
200
201 func TestReadMessageMaskedServerFrame(t *testing.T) {
202 mask := [4]byte{0x11, 0x22, 0x33, 0x44}
203 body := []byte("masked")
204 masked := []byte{:len(body)}
205 for i := int32(0); i < int32(len(body)); i++ {
206 masked[i] = body[i] ^ mask[i%4]
207 }
208 frame := []byte{:2 + 4 + int32(len(body))}
209 frame[0] = 0x80 | OpText
210 frame[1] = 0x80 | byte(len(body))
211 frame[2] = mask[0]
212 frame[3] = mask[1]
213 frame[4] = mask[2]
214 frame[5] = mask[3]
215 for i := int32(0); i < int32(len(masked)); i++ {
216 frame[6+i] = masked[i]
217 }
218
219 c, done := dialScript(t, frame)
220 if c == nil {
221 return
222 }
223 _, payload, err := c.ReadMessage()
224 if err != nil {
225 t.Fatalf("read: %s", err.Error())
226 }
227 if string(payload) != "masked" {
228 t.Fatalf("payload = %s", payload)
229 }
230 c.Close()
231 <-done
232 }
233
234 func TestReadMessageControlFrames(t *testing.T) {
235 // ping is answered with a pong and the following text frame is returned.
236 c, done := dialScript(t, rawFrame(OpPing, 3, nil, []byte("abc"))|rawFrame(OpText, 2, nil, []byte("ok")))
237 if c != nil {
238 op, payload, err := c.ReadMessage()
239 if err != nil {
240 t.Fatalf("ping then text: %s", err.Error())
241 }
242 if op != OpText || string(payload) != "ok" {
243 t.Fatalf("ping then text: op %d payload %s", int32(op), payload)
244 }
245 c.Close()
246 <-done
247 }
248
249 // pong is ignored, the following text frame is returned.
250 c2, done2 := dialScript(t, rawFrame(OpPong, 0, nil, nil)|rawFrame(OpText, 2, nil, []byte("go")))
251 if c2 != nil {
252 op2, payload2, err2 := c2.ReadMessage()
253 if err2 != nil {
254 t.Fatalf("pong then text: %s", err2.Error())
255 }
256 if op2 != OpText || string(payload2) != "go" {
257 t.Fatalf("pong then text: op %d payload %s", int32(op2), payload2)
258 }
259 c2.Close()
260 <-done2
261 }
262
263 // close is reported to the caller and answered.
264 c3, done3 := dialScript(t, rawFrame(OpClose, 2, nil, []byte{0x03, 0xE8}))
265 if c3 != nil {
266 op3, _, err3 := c3.ReadMessage()
267 if err3 != nil {
268 t.Fatalf("close: %s", err3.Error())
269 }
270 if op3 != OpClose {
271 t.Fatalf("close: op %d", int32(op3))
272 }
273 c3.Close()
274 <-done3
275 }
276 }
277
278 func TestReadMessageProtocolErrors(t *testing.T) {
279 // A reserved opcode is a protocol error, not a frame to skip.
280 c, done := dialScript(t, rawFrame(0x3, 1, nil, []byte("x")))
281 if c != nil {
282 if _, _, err := c.ReadMessage(); err == nil {
283 t.Fatal("a reserved opcode must fail")
284 }
285 c.Close()
286 <-done
287 }
288
289 // A declared payload above the limit is rejected before allocating.
290 over := int64(maxPayload) + 1
291 big := []byte{0, 0, 0, 0, byte(over >> 24), byte(over >> 16), byte(over >> 8), byte(over)}
292 c2, done2 := dialScript(t, rawFrame(OpText, 127, big, nil))
293 if c2 != nil {
294 if _, _, err2 := c2.ReadMessage(); err2 == nil {
295 t.Fatal("an oversized payload must fail")
296 }
297 c2.Close()
298 <-done2
299 }
300
301 // A frame that claims more bytes than it carries fails the full read.
302 c3, done3 := dialScript(t, rawFrame(OpText, 10, nil, []byte("xy")))
303 if c3 != nil {
304 if _, _, err3 := c3.ReadMessage(); err3 == nil {
305 t.Fatal("a truncated payload must fail")
306 }
307 c3.Close()
308 <-done3
309 }
310 }
311
312 func TestWriteClientFrameIsMasked(t *testing.T) {
313 fd, port, err := listenLoopback()
314 if err != nil {
315 t.Fatalf("listen: %v", err)
316 return
317 }
318 done := spawn(serveEchoFrame, fd, 2)
319 c, cerr := Dial("ws://127.0.0.1:" | portStr(port) | "/")
320 if cerr != nil {
321 t.Fatalf("dial: %v", cerr)
322 return
323 }
324
325 // Short payload and an extended-length payload: the server unmasks both and
326 // echoes them, so the client's own masking and header encoding are checked.
327 shortMsg := []byte("ping-pong")
328 if werr := c.WriteText(shortMsg); werr != nil {
329 t.Fatalf("write short: %s", werr.Error())
330 }
331 _, gotShort, rerr := c.ReadMessage()
332 if rerr != nil {
333 t.Fatalf("read short: %s", rerr.Error())
334 }
335 if string(gotShort) != string(shortMsg) {
336 t.Fatalf("short echo = %s", gotShort)
337 }
338
339 longMsg := makePayload(300)
340 if werr2 := c.WriteText(longMsg); werr2 != nil {
341 t.Fatalf("write long: %s", werr2.Error())
342 }
343 _, gotLong, rerr2 := c.ReadMessage()
344 if rerr2 != nil {
345 t.Fatalf("read long: %s", rerr2.Error())
346 }
347 if string(gotLong) != string(longMsg) {
348 t.Fatalf("long echo length %d, want %d", len(gotLong), len(longMsg))
349 }
350 c.Close()
351 <-done
352 }
353
354 func TestDialRejections(t *testing.T) {
355 // A URL that cannot parse.
356 if _, errA := Dial("://nope"); errA == nil {
357 t.Fatal("a malformed URL must fail")
358 }
359 // Connection refused: nothing listens on port 1.
360 if _, errB := Dial("ws://127.0.0.1:1/"); errB == nil {
361 t.Fatal("a refused dial must fail")
362 }
363 // Not an upgrade.
364 fd, port, err := listenLoopback()
365 if err != nil {
366 t.Fatalf("listen: %v", err)
367 return
368 }
369 done := spawn(serveReply, fd, []byte("HTTP/1.1 200 OK\r\n\r\n"))
370 if _, derr := Dial("ws://127.0.0.1:" | portStr(port) | "/"); derr == nil {
371 t.Fatal("a 200 reply must not upgrade")
372 }
373 <-done
374
375 // 101 with a wrong accept key.
376 fd2, port2, err2 := listenLoopback()
377 if err2 != nil {
378 t.Fatalf("listen: %v", err2)
379 return
380 }
381 bad := "HTTP/1.1 101 Switching Protocols\r\n" |
382 "Upgrade: websocket\r\n" |
383 "Connection: Upgrade\r\n" |
384 "Sec-WebSocket-Accept: not-the-key\r\n\r\n"
385 done2 := spawn(serveReply, fd2, []byte(bad))
386 if _, derr2 := Dial("ws://127.0.0.1:" | portStr(port2) | "/"); derr2 == nil {
387 t.Fatal("a wrong accept key must fail")
388 }
389 <-done2
390
391 // 101 with no accept header at all.
392 fd3, port3, err3 := listenLoopback()
393 if err3 != nil {
394 t.Fatalf("listen: %v", err3)
395 return
396 }
397 none := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\n\r\n"
398 done3 := spawn(serveReply, fd3, []byte(none))
399 if _, derr3 := Dial("ws://127.0.0.1:" | portStr(port3) | "/"); derr3 == nil {
400 t.Fatal("a missing accept header must fail")
401 }
402 <-done3
403 }
404
405 func TestComputeAcceptVector(t *testing.T) {
406 // RFC 6455 section 1.3: this key must produce this accept value.
407 got := ComputeAccept("dGhlIHNhbXBsZSBub25jZQ==")
408 if got != "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" {
409 t.Fatalf("ComputeAccept = %s", got)
410 }
411 }
412