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