package ws import ( "syscall" "testing" ) // Framing and handshake paths the subscribe test does not reach: both extended // length encodings, a masked server frame, ping/pong/close control frames, // protocol and truncation errors, the client's own masked write, and the // handshake rejections. // // Each case runs a loopback server in a spawn domain that writes a scripted // byte sequence after the upgrade, so no real relay and no TLS are involved. // serveScript completes the handshake and writes frames verbatim. func serveScript(fd int32, frames []byte) { nfd, _, err := syscall.Accept(fd) if err != nil { syscall.Close(fd) return } req := readUpgradeRequest(nfd) syscall.Write(nfd, []byte(upgradeHead(req))) if len(frames) > 0 { syscall.Write(nfd, frames) } // A short pause keeps the bytes in flight before the FIN; the client reads // what is buffered either way. syscall.Close(nfd) syscall.Close(fd) } // serveReply writes a raw HTTP reply to the upgrade request (for Dial errors). func serveReply(fd int32, reply []byte) { nfd, _, err := syscall.Accept(fd) if err != nil { syscall.Close(fd) return } readUpgradeRequest(nfd) syscall.Write(nfd, reply) syscall.Close(nfd) syscall.Close(fd) } // serveEchoFrame reads count masked client frames and echoes each payload as // an unmasked text frame, which proves the client masked it correctly. func serveEchoFrame(fd int32, count int32) { nfd, _, err := syscall.Accept(fd) if err != nil { syscall.Close(fd) return } req := readUpgradeRequest(nfd) syscall.Write(nfd, []byte(upgradeHead(req))) for i := int32(0); i < count; i++ { body := readClientFrame(nfd) syscall.Write(nfd, serverFrame(OpText, body)) } syscall.Close(nfd) syscall.Close(fd) } // readN reads exactly n bytes or returns what it got. func readN(fd int32, n int32) (out []byte) { out = []byte{:0:n} for len(out) < n { buf := []byte{:n - len(out)} got, err := syscall.Read(fd, buf) if got <= 0 || err != nil { return out } out = out | buf[:got] } return } // readClientFrame decodes one client frame (always masked) and unmasks it. func readClientFrame(fd int32) (payload []byte) { hdr := readN(fd, 2) if len(hdr) < 2 { return nil } plen := int32(hdr[1] & 0x7F) if plen == 126 { ext := readN(fd, 2) if len(ext) < 2 { return nil } plen = int32(ext[0])<<8 | int32(ext[1]) } else if plen == 127 { ext := readN(fd, 8) if len(ext) < 8 { return nil } plen = int32(ext[4])<<24 | int32(ext[5])<<16 | int32(ext[6])<<8 | int32(ext[7]) } var mask [4]byte if hdr[1]&0x80 != 0 { m := readN(fd, 4) if len(m) < 4 { return nil } mask = [4]byte{m[0], m[1], m[2], m[3]} } body := readN(fd, plen) for i := int32(0); i < int32(len(body)); i++ { body[i] = body[i] ^ mask[i%4] } return body } // rawFrame builds an unmasked frame with an explicit length byte, so a test can // exercise the 126 and 127 encodings and the masked variant. func rawFrame(op byte, plen int32, extended []byte, payload []byte) (buf []byte) { buf = []byte{:2} buf[0] = 0x80 | op buf[1] = byte(plen) buf = buf | extended return buf | payload } func dialScript(t *testing.T, frames []byte) (c *Conn, done chan struct{}) { t.Helper() fd, port, err := listenLoopback() if err != nil { t.Fatalf("listen: %v", err) return } done = spawn(serveScript, fd, frames) c, cerr := Dial("ws://127.0.0.1:" | portStr(port) | "/") if cerr != nil { t.Fatalf("dial: %v", cerr) return } return } func portStr(port int32) (s string) { if port == 0 { return "0" } buf := []byte{:6} n := int32(6) for port > 0 { n-- buf[n] = byte('0' + port%10) port = port / 10 } return string(buf[n:]) } func TestReadMessageLengthEncodings(t *testing.T) { // 126: 200 bytes, extended 16-bit length. small := makePayload(200) cases := []struct { name string op byte frame []byte want []byte }{ {"short", OpText, rawFrame(OpText, 5, nil, []byte("hello")), []byte("hello")}, {"extended16", OpText, rawFrame(OpText, 126, []byte{byte(200 >> 8), byte(200)}, small), small}, {"extended64", OpText, rawFrame(OpText, 127, []byte{0, 0, 0, 0, byte(70000 >> 24), byte(70000 >> 16), byte(70000 >> 8), byte(70000)}, makePayload(70000)), makePayload(70000)}, {"binary", OpBinary, rawFrame(OpBinary, 3, nil, []byte{1, 2, 3}), []byte{1, 2, 3}}, {"empty", OpText, rawFrame(OpText, 0, nil, nil), nil}, } for _, tc := range cases { c, done := dialScript(t, tc.frame) if c == nil { continue } op, payload, err := c.ReadMessage() if err != nil { t.Fatalf("%s: %s", tc.name, err.Error()) } if op != tc.op { t.Fatalf("%s: op = %d, want %d", tc.name, int32(op), int32(tc.op)) } if string(payload) != string(tc.want) { t.Fatalf("%s: payload length %d, want %d", tc.name, len(payload), len(tc.want)) } c.Close() <-done } } func makePayload(n int32) (b []byte) { b = []byte{:n} for i := int32(0); i < n; i++ { b[i] = byte('a' + i%26) } return } func TestReadMessageMaskedServerFrame(t *testing.T) { mask := [4]byte{0x11, 0x22, 0x33, 0x44} body := []byte("masked") masked := []byte{:len(body)} for i := int32(0); i < int32(len(body)); i++ { masked[i] = body[i] ^ mask[i%4] } frame := []byte{:2 + 4 + int32(len(body))} frame[0] = 0x80 | OpText frame[1] = 0x80 | byte(len(body)) frame[2] = mask[0] frame[3] = mask[1] frame[4] = mask[2] frame[5] = mask[3] for i := int32(0); i < int32(len(masked)); i++ { frame[6+i] = masked[i] } c, done := dialScript(t, frame) if c == nil { return } _, payload, err := c.ReadMessage() if err != nil { t.Fatalf("read: %s", err.Error()) } if string(payload) != "masked" { t.Fatalf("payload = %s", payload) } c.Close() <-done } func TestReadMessageControlFrames(t *testing.T) { // ping is answered with a pong and the following text frame is returned. c, done := dialScript(t, rawFrame(OpPing, 3, nil, []byte("abc"))|rawFrame(OpText, 2, nil, []byte("ok"))) if c != nil { op, payload, err := c.ReadMessage() if err != nil { t.Fatalf("ping then text: %s", err.Error()) } if op != OpText || string(payload) != "ok" { t.Fatalf("ping then text: op %d payload %s", int32(op), payload) } c.Close() <-done } // pong is ignored, the following text frame is returned. c2, done2 := dialScript(t, rawFrame(OpPong, 0, nil, nil)|rawFrame(OpText, 2, nil, []byte("go"))) if c2 != nil { op2, payload2, err2 := c2.ReadMessage() if err2 != nil { t.Fatalf("pong then text: %s", err2.Error()) } if op2 != OpText || string(payload2) != "go" { t.Fatalf("pong then text: op %d payload %s", int32(op2), payload2) } c2.Close() <-done2 } // close is reported to the caller and answered. c3, done3 := dialScript(t, rawFrame(OpClose, 2, nil, []byte{0x03, 0xE8})) if c3 != nil { op3, _, err3 := c3.ReadMessage() if err3 != nil { t.Fatalf("close: %s", err3.Error()) } if op3 != OpClose { t.Fatalf("close: op %d", int32(op3)) } c3.Close() <-done3 } } func TestReadMessageProtocolErrors(t *testing.T) { // A reserved opcode is a protocol error, not a frame to skip. c, done := dialScript(t, rawFrame(0x3, 1, nil, []byte("x"))) if c != nil { if _, _, err := c.ReadMessage(); err == nil { t.Fatal("a reserved opcode must fail") } c.Close() <-done } // A declared payload above the limit is rejected before allocating. over := int64(maxPayload) + 1 big := []byte{0, 0, 0, 0, byte(over >> 24), byte(over >> 16), byte(over >> 8), byte(over)} c2, done2 := dialScript(t, rawFrame(OpText, 127, big, nil)) if c2 != nil { if _, _, err2 := c2.ReadMessage(); err2 == nil { t.Fatal("an oversized payload must fail") } c2.Close() <-done2 } // A frame that claims more bytes than it carries fails the full read. c3, done3 := dialScript(t, rawFrame(OpText, 10, nil, []byte("xy"))) if c3 != nil { if _, _, err3 := c3.ReadMessage(); err3 == nil { t.Fatal("a truncated payload must fail") } c3.Close() <-done3 } } func TestWriteClientFrameIsMasked(t *testing.T) { fd, port, err := listenLoopback() if err != nil { t.Fatalf("listen: %v", err) return } done := spawn(serveEchoFrame, fd, 2) c, cerr := Dial("ws://127.0.0.1:" | portStr(port) | "/") if cerr != nil { t.Fatalf("dial: %v", cerr) return } // Short payload and an extended-length payload: the server unmasks both and // echoes them, so the client's own masking and header encoding are checked. shortMsg := []byte("ping-pong") if werr := c.WriteText(shortMsg); werr != nil { t.Fatalf("write short: %s", werr.Error()) } _, gotShort, rerr := c.ReadMessage() if rerr != nil { t.Fatalf("read short: %s", rerr.Error()) } if string(gotShort) != string(shortMsg) { t.Fatalf("short echo = %s", gotShort) } longMsg := makePayload(300) if werr2 := c.WriteText(longMsg); werr2 != nil { t.Fatalf("write long: %s", werr2.Error()) } _, gotLong, rerr2 := c.ReadMessage() if rerr2 != nil { t.Fatalf("read long: %s", rerr2.Error()) } if string(gotLong) != string(longMsg) { t.Fatalf("long echo length %d, want %d", len(gotLong), len(longMsg)) } c.Close() <-done } func TestDialRejections(t *testing.T) { // A URL that cannot parse. if _, errA := Dial("://nope"); errA == nil { t.Fatal("a malformed URL must fail") } // Connection refused: nothing listens on port 1. if _, errB := Dial("ws://127.0.0.1:1/"); errB == nil { t.Fatal("a refused dial must fail") } // Not an upgrade. fd, port, err := listenLoopback() if err != nil { t.Fatalf("listen: %v", err) return } done := spawn(serveReply, fd, []byte("HTTP/1.1 200 OK\r\n\r\n")) if _, derr := Dial("ws://127.0.0.1:" | portStr(port) | "/"); derr == nil { t.Fatal("a 200 reply must not upgrade") } <-done // 101 with a wrong accept key. fd2, port2, err2 := listenLoopback() if err2 != nil { t.Fatalf("listen: %v", err2) return } bad := "HTTP/1.1 101 Switching Protocols\r\n" | "Upgrade: websocket\r\n" | "Connection: Upgrade\r\n" | "Sec-WebSocket-Accept: not-the-key\r\n\r\n" done2 := spawn(serveReply, fd2, []byte(bad)) if _, derr2 := Dial("ws://127.0.0.1:" | portStr(port2) | "/"); derr2 == nil { t.Fatal("a wrong accept key must fail") } <-done2 // 101 with no accept header at all. fd3, port3, err3 := listenLoopback() if err3 != nil { t.Fatalf("listen: %v", err3) return } none := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\n\r\n" done3 := spawn(serveReply, fd3, []byte(none)) if _, derr3 := Dial("ws://127.0.0.1:" | portStr(port3) | "/"); derr3 == nil { t.Fatal("a missing accept header must fail") } <-done3 } func TestComputeAcceptVector(t *testing.T) { // RFC 6455 section 1.3: this key must produce this accept value. got := ComputeAccept("dGhlIHNhbXBsZSBub25jZQ==") if got != "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" { t.Fatalf("ComputeAccept = %s", got) } }