Cover WebSocket session lifecycle (result channels, SNMP batch flushing, heartbeats, restart), MikroTik TLS TOFU verification mismatch, worker pool exhaustion, plaintext WebSocket rejection, HTTP/TCP/DNS check execution, IPv4 fallback, and customer network edge cases. Make heartbeat intervals and toWebSocketURL exit injectable for testing.
838 lines
20 KiB
Go
838 lines
20 KiB
Go
package main
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"encoding/binary"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"net/http"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestWriteFrameMasked(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
|
|
payload := []byte("hello")
|
|
if err := ws.writeFrame(opText, payload); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
frame := buf.Bytes()
|
|
|
|
if frame[0] != 0x81 {
|
|
t.Errorf("first byte: got %#x, want 0x81", frame[0])
|
|
}
|
|
|
|
if frame[1] != 0x85 {
|
|
t.Errorf("second byte: got %#x, want 0x85", frame[1])
|
|
}
|
|
|
|
maskKey := frame[2:6]
|
|
maskedPayload := frame[6:]
|
|
|
|
for i := range maskedPayload {
|
|
maskedPayload[i] ^= maskKey[i%4]
|
|
}
|
|
if string(maskedPayload) != "hello" {
|
|
t.Errorf("unmasked payload: got %q, want %q", maskedPayload, "hello")
|
|
}
|
|
}
|
|
|
|
func TestWriteFrame16BitLength(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
|
|
payload := make([]byte, 300)
|
|
if err := ws.writeFrame(opBinary, payload); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
frame := buf.Bytes()
|
|
if frame[1]&0x7F != 126 {
|
|
t.Errorf("expected 126 length marker, got %d", frame[1]&0x7F)
|
|
}
|
|
extLen := binary.BigEndian.Uint16(frame[2:4])
|
|
if extLen != 300 {
|
|
t.Errorf("extended length: got %d, want 300", extLen)
|
|
}
|
|
}
|
|
|
|
func TestWriteFrame64BitLength(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
|
|
payload := make([]byte, 70000) // > 65535, uses 8-byte extended
|
|
if err := ws.writeFrame(opBinary, payload); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
frame := buf.Bytes()
|
|
if frame[1]&0x7F != 127 {
|
|
t.Errorf("expected 127 length marker, got %d", frame[1]&0x7F)
|
|
}
|
|
extLen := binary.BigEndian.Uint64(frame[2:10])
|
|
if extLen != 70000 {
|
|
t.Errorf("extended length: got %d, want 70000", extLen)
|
|
}
|
|
}
|
|
|
|
func TestReadFrame(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
payload := []byte("world")
|
|
buf.WriteByte(0x81) // FIN + text
|
|
buf.WriteByte(byte(len(payload)))
|
|
buf.Write(payload)
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
opcode, data, err := ws.readFrame()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if opcode != opText {
|
|
t.Errorf("opcode: got %d, want %d", opcode, opText)
|
|
}
|
|
if string(data) != "world" {
|
|
t.Errorf("data: got %q, want %q", data, "world")
|
|
}
|
|
}
|
|
|
|
func TestReadFrame16BitLength(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
payload := make([]byte, 300)
|
|
for i := range payload {
|
|
payload[i] = byte(i % 256)
|
|
}
|
|
buf.WriteByte(0x82) // FIN + binary
|
|
buf.WriteByte(126) // 16-bit extended length
|
|
var extLen [2]byte
|
|
binary.BigEndian.PutUint16(extLen[:], 300)
|
|
buf.Write(extLen[:])
|
|
buf.Write(payload)
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
opcode, data, err := ws.readFrame()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if opcode != opBinary {
|
|
t.Errorf("opcode: got %d, want %d", opcode, opBinary)
|
|
}
|
|
if len(data) != 300 {
|
|
t.Errorf("data length: got %d, want 300", len(data))
|
|
}
|
|
}
|
|
|
|
func TestReadFrame64BitLength(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
payload := make([]byte, 70000)
|
|
buf.WriteByte(0x82) // FIN + binary
|
|
buf.WriteByte(127) // 64-bit extended length
|
|
var extLen [8]byte
|
|
binary.BigEndian.PutUint64(extLen[:], 70000)
|
|
buf.Write(extLen[:])
|
|
buf.Write(payload)
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
opcode, data, err := ws.readFrame()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if opcode != opBinary {
|
|
t.Errorf("opcode: got %d, want %d", opcode, opBinary)
|
|
}
|
|
if len(data) != 70000 {
|
|
t.Errorf("data length: got %d, want 70000", len(data))
|
|
}
|
|
}
|
|
|
|
func TestReadFrameMasked(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
payload := []byte("test")
|
|
maskKey := [4]byte{0x12, 0x34, 0x56, 0x78}
|
|
masked := make([]byte, len(payload))
|
|
for i := range payload {
|
|
masked[i] = payload[i] ^ maskKey[i%4]
|
|
}
|
|
|
|
buf.WriteByte(0x81) // FIN + text
|
|
buf.WriteByte(0x80 | byte(len(payload))) // masked + length
|
|
buf.Write(maskKey[:])
|
|
buf.Write(masked)
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
opcode, data, err := ws.readFrame()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if opcode != opText {
|
|
t.Errorf("opcode: got %d, want %d", opcode, opText)
|
|
}
|
|
if string(data) != "test" {
|
|
t.Errorf("data: got %q, want %q", data, "test")
|
|
}
|
|
}
|
|
|
|
func TestReadMessageClose(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
// Close frame
|
|
buf.WriteByte(0x80 | byte(opClose))
|
|
buf.WriteByte(0) // no payload
|
|
|
|
rw := &captureWriter{Reader: &buf}
|
|
ws := testWSConn(&nopCloser{readWriter: rw})
|
|
|
|
_, opcode, err := ws.ReadMessage()
|
|
if err != io.EOF {
|
|
t.Errorf("expected io.EOF, got %v", err)
|
|
}
|
|
if opcode != opClose {
|
|
t.Errorf("opcode: got %d, want %d", opcode, opClose)
|
|
}
|
|
// Verify close reply was sent
|
|
if len(rw.written) == 0 {
|
|
t.Error("expected close reply frame to be written")
|
|
}
|
|
}
|
|
|
|
func TestReadFramePingPong(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
|
|
// Ping frame
|
|
buf.WriteByte(0x80 | byte(opPing))
|
|
buf.WriteByte(0) // no payload
|
|
|
|
// Then a text frame
|
|
text := []byte("data")
|
|
buf.WriteByte(0x81)
|
|
buf.WriteByte(byte(len(text)))
|
|
buf.Write(text)
|
|
|
|
rw := &captureWriter{Reader: &buf}
|
|
ws := testWSConn(&nopCloser{readWriter: rw})
|
|
|
|
data, _, err := ws.ReadMessage()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if string(data) != "data" {
|
|
t.Errorf("got %q, want %q", data, "data")
|
|
}
|
|
|
|
if len(rw.written) == 0 {
|
|
t.Error("expected pong frame to be written")
|
|
}
|
|
}
|
|
|
|
func TestReadFrame16BitLengthError(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
buf.WriteByte(0x81) // FIN + text
|
|
buf.WriteByte(126) // 16-bit length indicator, but no length bytes follow
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
_, _, err := ws.readFrame()
|
|
if err == nil {
|
|
t.Error("expected error for truncated 16-bit length")
|
|
}
|
|
}
|
|
|
|
func TestReadFrame64BitLengthError(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
buf.WriteByte(0x81) // FIN + text
|
|
buf.WriteByte(127) // 64-bit length indicator, but no length bytes follow
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
_, _, err := ws.readFrame()
|
|
if err == nil {
|
|
t.Error("expected error for truncated 64-bit length")
|
|
}
|
|
}
|
|
|
|
func TestReadFrameMaskKeyError(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
buf.WriteByte(0x81) // FIN + text
|
|
buf.WriteByte(0x80 | 0x01) // masked + length 1, but no mask key or payload
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
_, _, err := ws.readFrame()
|
|
if err == nil {
|
|
t.Error("expected error for missing mask key")
|
|
}
|
|
}
|
|
|
|
func TestReadFramePayloadError(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
buf.WriteByte(0x81) // FIN + text
|
|
buf.WriteByte(5) // length 5, but no payload
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
_, _, err := ws.readFrame()
|
|
if err == nil {
|
|
t.Error("expected error for truncated payload")
|
|
}
|
|
}
|
|
|
|
func TestWriteFrameHeaderError(t *testing.T) {
|
|
ws := testWSConn(&failOnWriteBuffer{Buffer: bytes.NewBuffer(nil)})
|
|
err := ws.writeFrame(opText, []byte("hello"))
|
|
if err == nil {
|
|
t.Error("expected write error")
|
|
}
|
|
}
|
|
|
|
func TestReadMessageError(t *testing.T) {
|
|
// Empty buffer causes immediate EOF on readFrame
|
|
buf := bytes.NewBuffer(nil)
|
|
ws := testWSConn(&nopCloser{readWriter: buf})
|
|
_, _, err := ws.ReadMessage()
|
|
if err == nil {
|
|
t.Error("expected error for empty buffer")
|
|
}
|
|
}
|
|
|
|
func TestReadMessagePongError(t *testing.T) {
|
|
// Ping frame followed by write failure for pong
|
|
var buf bytes.Buffer
|
|
buf.WriteByte(0x80 | byte(opPing))
|
|
buf.WriteByte(0)
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &failOnWriteBuffer{Buffer: &buf}})
|
|
_, _, err := ws.ReadMessage()
|
|
if err == nil {
|
|
t.Error("expected pong write error")
|
|
}
|
|
}
|
|
|
|
func TestReadFrameExceedsMaxSize(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
buf.WriteByte(0x82) // FIN + binary
|
|
buf.WriteByte(127) // 64-bit extended length
|
|
var extLen [8]byte
|
|
binary.BigEndian.PutUint64(extLen[:], uint64(maxFrameSize+1))
|
|
buf.Write(extLen[:])
|
|
// Don't need to write the payload — should reject before reading it
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
_, _, err := ws.readFrame()
|
|
if err == nil {
|
|
t.Error("expected error for frame exceeding max size")
|
|
}
|
|
if !strings.Contains(err.Error(), "exceeds max") {
|
|
t.Errorf("expected 'exceeds max' in error, got: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestReadFrameEmptyPayload(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
buf.WriteByte(0x82) // FIN + binary
|
|
buf.WriteByte(0) // zero length
|
|
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
opcode, data, err := ws.readFrame()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if opcode != opBinary {
|
|
t.Errorf("opcode: got %d, want %d", opcode, opBinary)
|
|
}
|
|
if len(data) != 0 {
|
|
t.Errorf("expected empty payload, got %d bytes", len(data))
|
|
}
|
|
}
|
|
|
|
func TestWriteFrameEmptyPayload(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
|
|
if err := ws.writeFrame(opClose, nil); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
frame := buf.Bytes()
|
|
if frame[0]&0x0F != opClose {
|
|
t.Errorf("expected close opcode, got %d", frame[0]&0x0F)
|
|
}
|
|
// Length should be 0 (masked)
|
|
if frame[1]&0x7F != 0 {
|
|
t.Errorf("expected 0 length, got %d", frame[1]&0x7F)
|
|
}
|
|
}
|
|
|
|
// failOnWriteBuffer reads from Buffer but fails on Write.
|
|
type failOnWriteBuffer struct {
|
|
*bytes.Buffer
|
|
}
|
|
|
|
func (f *failOnWriteBuffer) Write(p []byte) (int, error) {
|
|
return 0, fmt.Errorf("write failed")
|
|
}
|
|
func (f *failOnWriteBuffer) Close() error { return nil }
|
|
|
|
func TestWriteText(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
|
|
if err := ws.WriteText([]byte("hello")); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
// Verify it wrote a text frame (opcode 1)
|
|
frame := buf.Bytes()
|
|
if frame[0]&0x0F != opText {
|
|
t.Errorf("expected text opcode, got %d", frame[0]&0x0F)
|
|
}
|
|
}
|
|
|
|
func TestClose(t *testing.T) {
|
|
var buf bytes.Buffer
|
|
ws := testWSConn(&nopCloser{readWriter: &buf})
|
|
|
|
if err := ws.Close(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
// Verify it wrote a close frame (opcode 8)
|
|
frame := buf.Bytes()
|
|
if len(frame) == 0 {
|
|
t.Fatal("expected close frame to be written")
|
|
}
|
|
if frame[0]&0x0F != opClose {
|
|
t.Errorf("expected close opcode, got %d", frame[0]&0x0F)
|
|
}
|
|
}
|
|
|
|
func TestWSDial(t *testing.T) {
|
|
// Start a test TCP server that responds with a valid WebSocket upgrade
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
go func() {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
return
|
|
}
|
|
defer func() { _ = conn.Close() }()
|
|
|
|
buf := make([]byte, 4096)
|
|
n, _ := conn.Read(buf)
|
|
reqStr := string(buf[:n])
|
|
|
|
// Extract Sec-WebSocket-Key from request
|
|
key := extractHeader(reqStr, "Sec-WebSocket-Key")
|
|
accept := computeAcceptKey(key)
|
|
|
|
resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: " + accept + "\r\n\r\n"
|
|
_, _ = conn.Write([]byte(resp))
|
|
|
|
// Keep connection open briefly for the test
|
|
buf2 := make([]byte, 1)
|
|
_, _ = conn.Read(buf2)
|
|
}()
|
|
|
|
addr := ln.Addr().String()
|
|
ws, err := WSDial("ws://" + addr + "/socket")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
_ = ws.Close()
|
|
}
|
|
|
|
func TestWSDialHandshakeFailed(t *testing.T) {
|
|
// Start a test HTTP server that returns 403
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
go func() {
|
|
for {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
return
|
|
}
|
|
go func(c net.Conn) {
|
|
defer func() { _ = c.Close() }()
|
|
buf := make([]byte, 4096)
|
|
_, _ = c.Read(buf)
|
|
resp := "HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\n\r\n"
|
|
_, _ = c.Write([]byte(resp))
|
|
}(conn)
|
|
}
|
|
}()
|
|
|
|
addr := ln.Addr().String()
|
|
_, err = WSDial("ws://" + addr + "/socket")
|
|
if err == nil {
|
|
t.Error("expected handshake error")
|
|
}
|
|
if !strings.Contains(err.Error(), "handshake failed") {
|
|
t.Errorf("expected 'handshake failed' in error, got: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWSDialReadHandshakeError(t *testing.T) {
|
|
// Server accepts connection but immediately closes without sending response
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
go func() {
|
|
for {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
return
|
|
}
|
|
go func(c net.Conn) {
|
|
buf := make([]byte, 4096)
|
|
_, _ = c.Read(buf)
|
|
_ = c.Close()
|
|
}(conn)
|
|
}
|
|
}()
|
|
|
|
addr := ln.Addr().String()
|
|
_, err = WSDial("ws://" + addr + "/socket")
|
|
if err == nil {
|
|
t.Error("expected read handshake error")
|
|
}
|
|
}
|
|
|
|
func TestWSDialGenerateKeyError(t *testing.T) {
|
|
origRandRead := randRead
|
|
defer func() { randRead = origRandRead }()
|
|
randRead = func(b []byte) (int, error) {
|
|
return 0, fmt.Errorf("entropy exhausted")
|
|
}
|
|
|
|
// Start a TCP server that accepts connections
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
go func() {
|
|
conn, _ := ln.Accept()
|
|
if conn != nil {
|
|
_ = conn.Close()
|
|
}
|
|
}()
|
|
|
|
addr := ln.Addr().String()
|
|
_, err = WSDial("ws://" + addr + "/socket")
|
|
if err == nil {
|
|
t.Error("expected generate key error")
|
|
}
|
|
if !strings.Contains(err.Error(), "generate key") {
|
|
t.Errorf("expected 'generate key' in error, got: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWSDialWriteHandshakeError(t *testing.T) {
|
|
origDial := netDial
|
|
defer func() { netDial = origDial }()
|
|
|
|
netDial = func(network, addr string) (net.Conn, error) {
|
|
// Return a connection whose Write always fails
|
|
client, _ := net.Pipe()
|
|
_ = client.Close() // Close immediately so Write fails
|
|
return client, nil
|
|
}
|
|
|
|
_, err := WSDial("ws://127.0.0.1:9999/socket")
|
|
if err == nil {
|
|
t.Error("expected write handshake error")
|
|
}
|
|
if !strings.Contains(err.Error(), "write handshake") {
|
|
t.Errorf("expected 'write handshake' in error, got: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWSDialBadURL(t *testing.T) {
|
|
_, err := WSDial("://bad url")
|
|
if err == nil {
|
|
t.Error("expected parse error for bad URL")
|
|
}
|
|
}
|
|
|
|
func TestWSDialHandshakeTimeout(t *testing.T) {
|
|
// Server accepts connection but never sends data — should trigger timeout
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
go func() {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
return
|
|
}
|
|
// Read the request but never respond
|
|
buf := make([]byte, 4096)
|
|
_, _ = conn.Read(buf)
|
|
// Hold connection open until test ends
|
|
<-make(chan struct{})
|
|
_ = conn.Close()
|
|
}()
|
|
|
|
origTimeout := wsHandshakeTimeout
|
|
defer func() { wsHandshakeTimeout = origTimeout }()
|
|
wsHandshakeTimeout = 500 * time.Millisecond // Short timeout for test
|
|
|
|
addr := ln.Addr().String()
|
|
start := time.Now()
|
|
_, err = WSDial("ws://" + addr + "/socket")
|
|
elapsed := time.Since(start)
|
|
|
|
if err == nil {
|
|
t.Error("expected timeout error")
|
|
}
|
|
if elapsed > 5*time.Second {
|
|
t.Errorf("took too long (%v), timeout didn't work", elapsed)
|
|
}
|
|
}
|
|
|
|
func TestWSDialIPv4Fallback(t *testing.T) {
|
|
// Start a test TCP server on IPv4 that responds with valid WebSocket upgrade
|
|
ln, err := net.Listen("tcp4", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
go func() {
|
|
for {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
return
|
|
}
|
|
go func(c net.Conn) {
|
|
defer func() { _ = c.Close() }()
|
|
buf := make([]byte, 4096)
|
|
n, _ := c.Read(buf)
|
|
key := extractHeader(string(buf[:n]), "Sec-WebSocket-Key")
|
|
accept := computeAcceptKey(key)
|
|
resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: " + accept + "\r\n\r\n"
|
|
_, _ = c.Write([]byte(resp))
|
|
b := make([]byte, 1)
|
|
_, _ = c.Read(b)
|
|
}(conn)
|
|
}
|
|
}()
|
|
|
|
addr := ln.Addr().String()
|
|
callCount := 0
|
|
origDial := netDial
|
|
defer func() { netDial = origDial }()
|
|
|
|
netDial = func(network, a string) (net.Conn, error) {
|
|
callCount++
|
|
if network == "tcp" {
|
|
// Simulate IPv6 failure: TCP connects but server rejects
|
|
return nil, fmt.Errorf("dial tcp [::1]:%s: connect: connection refused", a)
|
|
}
|
|
// tcp4 fallback goes to the real server
|
|
return net.Dial(network, addr)
|
|
}
|
|
|
|
ws, err := WSDial("ws://" + addr + "/socket")
|
|
if err != nil {
|
|
t.Fatalf("expected IPv4 fallback to succeed: %v", err)
|
|
}
|
|
_ = ws.Close()
|
|
|
|
if callCount < 2 {
|
|
t.Errorf("expected at least 2 dial attempts (tcp + tcp4), got %d", callCount)
|
|
}
|
|
}
|
|
|
|
func TestWSDialTLSPath(t *testing.T) {
|
|
// Verify the TLS dial path is exercised (connection will fail, but we hit the code path)
|
|
origTLS := tlsDial
|
|
defer func() { tlsDial = origTLS }()
|
|
|
|
called := false
|
|
tlsDial = func(network, addr string) (net.Conn, error) {
|
|
called = true
|
|
return nil, fmt.Errorf("tls dial: connection refused")
|
|
}
|
|
|
|
_, err := WSDial("wss://127.0.0.1:9999/path")
|
|
if err == nil {
|
|
t.Error("expected TLS dial error")
|
|
}
|
|
if !called {
|
|
t.Error("expected tlsDial to be called for wss:// URL")
|
|
}
|
|
}
|
|
|
|
func TestWSDialMissingAcceptHeader(t *testing.T) {
|
|
// Server sends 101 but without Sec-WebSocket-Accept header
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
go func() {
|
|
for {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
return
|
|
}
|
|
go func(c net.Conn) {
|
|
defer func() { _ = c.Close() }()
|
|
buf := make([]byte, 4096)
|
|
_, _ = c.Read(buf)
|
|
resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n"
|
|
_, _ = c.Write([]byte(resp))
|
|
}(conn)
|
|
}
|
|
}()
|
|
|
|
addr := ln.Addr().String()
|
|
_, err = WSDial("ws://" + addr + "/socket")
|
|
if err == nil {
|
|
t.Error("expected error for missing accept header")
|
|
}
|
|
if !strings.Contains(err.Error(), "missing Sec-WebSocket-Accept") {
|
|
t.Errorf("expected 'missing Sec-WebSocket-Accept' in error, got: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWSDialDefaultPorts(t *testing.T) {
|
|
// Test that ws:// defaults to port 80 — will fail to connect but verifies URL parsing
|
|
_, err := WSDial("ws://127.0.0.1/path")
|
|
if err == nil {
|
|
t.Error("expected connection error (nothing on port 80)")
|
|
}
|
|
|
|
// Test that wss:// defaults to port 443
|
|
_, err = WSDial("wss://127.0.0.1/path")
|
|
if err == nil {
|
|
t.Error("expected connection error (nothing on port 443)")
|
|
}
|
|
}
|
|
|
|
func TestWSDialRealHTTPServer(t *testing.T) {
|
|
// Test with a real HTTP server that sends proper 101 upgrade
|
|
mux := http.NewServeMux()
|
|
mux.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
|
|
hj, ok := w.(http.Hijacker)
|
|
if !ok {
|
|
http.Error(w, "hijack not supported", 500)
|
|
return
|
|
}
|
|
key := r.Header.Get("Sec-WebSocket-Key")
|
|
accept := computeAcceptKey(key)
|
|
conn, brw, _ := hj.Hijack()
|
|
defer func() { _ = conn.Close() }()
|
|
_, _ = brw.WriteString("HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: " + accept + "\r\n\r\n")
|
|
_ = brw.Flush()
|
|
// Keep alive briefly
|
|
buf := make([]byte, 1)
|
|
_, _ = conn.Read(buf)
|
|
})
|
|
ln, _ := net.Listen("tcp", "127.0.0.1:0")
|
|
defer func() { _ = ln.Close() }()
|
|
go func() { _ = http.Serve(ln, mux) }()
|
|
|
|
addr := ln.Addr().String()
|
|
ws, err := WSDial("ws://" + addr + "/ws")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
_ = ws.Close()
|
|
}
|
|
|
|
func TestComputeAcceptKey(t *testing.T) {
|
|
// RFC 6455 Section 4.2.2 test vector
|
|
got := computeAcceptKey("dGhlIHNhbXBsZSBub25jZQ==")
|
|
want := "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="
|
|
if got != want {
|
|
t.Errorf("computeAcceptKey = %q, want %q", got, want)
|
|
}
|
|
}
|
|
|
|
func TestWSDialRejectsInvalidAcceptKey(t *testing.T) {
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
go func() {
|
|
for {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
return
|
|
}
|
|
go func(c net.Conn) {
|
|
defer func() { _ = c.Close() }()
|
|
buf := make([]byte, 4096)
|
|
_, _ = c.Read(buf)
|
|
resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: INVALID_KEY\r\n\r\n"
|
|
_, _ = c.Write([]byte(resp))
|
|
}(conn)
|
|
}
|
|
}()
|
|
|
|
addr := ln.Addr().String()
|
|
_, err = WSDial("ws://" + addr + "/socket")
|
|
if err == nil {
|
|
t.Error("expected error for invalid accept key")
|
|
}
|
|
if !strings.Contains(err.Error(), "accept key") {
|
|
t.Errorf("expected 'accept key' in error, got: %v", err)
|
|
}
|
|
}
|
|
|
|
func testWSConn(rw io.ReadWriteCloser) *WSConn {
|
|
return &WSConn{conn: rw, reader: bufio.NewReader(rw)}
|
|
}
|
|
|
|
// nopCloser wraps a ReadWriter with a no-op Close.
|
|
type nopCloser struct {
|
|
readWriter interface {
|
|
Read([]byte) (int, error)
|
|
Write([]byte) (int, error)
|
|
}
|
|
}
|
|
|
|
func (n *nopCloser) Read(p []byte) (int, error) { return n.readWriter.Read(p) }
|
|
func (n *nopCloser) Write(p []byte) (int, error) { return n.readWriter.Write(p) }
|
|
func (n *nopCloser) Close() error { return nil }
|
|
|
|
// extractHeader extracts a header value from a raw HTTP request string.
|
|
func extractHeader(req, name string) string {
|
|
for _, line := range strings.Split(req, "\r\n") {
|
|
if strings.HasPrefix(strings.ToLower(line), strings.ToLower(name)+": ") {
|
|
return strings.TrimSpace(line[len(name)+2:])
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
|
|
// captureWriter captures written data while reading from a separate Reader.
|
|
type captureWriter struct {
|
|
Reader *bytes.Buffer
|
|
written []byte
|
|
}
|
|
|
|
func (c *captureWriter) Read(p []byte) (int, error) { return c.Reader.Read(p) }
|
|
func (c *captureWriter) Write(p []byte) (int, error) {
|
|
c.written = append(c.written, p...)
|
|
return len(p), nil
|
|
}
|
|
func (c *captureWriter) Close() error { return nil }
|