fix: 测试 WS 客户端按字节流拆包并恢复群事件同步下发

This commit is contained in:
Nixevol
2026-09-30 16:24:31 +08:00
parent 5f91a7758b
commit 09fb544b7d
6 changed files with 265 additions and 67 deletions
+39 -8
View File
@@ -12,6 +12,7 @@ import (
"net/http"
"net/url"
"strings"
"sync"
"time"
)
@@ -164,25 +165,35 @@ func (c *tcpMQTT) RemoteAddr() net.Addr {
type wsMQTT struct {
conn net.Conn
r *bufio.Reader
buf []byte
wmu sync.Mutex
}
func (c *wsMQTT) Send(packet []byte) error {
c.wmu.Lock()
defer c.wmu.Unlock()
return writeWSClientBinary(c.conn, packet)
}
func (c *wsMQTT) Recv() ([]byte, error) {
for {
if pkt, n, ok := splitMQTTPacket(c.buf); ok {
c.buf = c.buf[n:]
return pkt, nil
}
payload, opcode, err := readWSFrame(c.r)
if err != nil {
return nil, err
}
switch opcode {
case 0x2:
return payload, nil
case 0x0, 0x2:
c.buf = append(c.buf, payload...)
case 0x8:
return nil, io.EOF
case 0x9:
c.wmu.Lock()
_ = writeWSClientControl(c.conn, 0xA, payload)
c.wmu.Unlock()
case 0xA:
continue
default:
@@ -265,17 +276,37 @@ func writeWSClientFrame(w io.Writer, opcode byte, payload []byte) error {
header = append(header, ext[:]...)
}
header = append(header, mask...)
masked := make([]byte, n)
frame := make([]byte, len(header)+n)
copy(frame, header)
for i := 0; i < n; i++ {
masked[i] = payload[i] ^ mask[i%4]
frame[len(header)+i] = payload[i] ^ mask[i%4]
}
if _, err := w.Write(header); err != nil {
return err
}
_, err := w.Write(masked)
_, err := w.Write(frame)
return err
}
// splitMQTTPacket 从缓冲头部切出一个完整 MQTT 控制包。不够一包时返回 ok=false。
func splitMQTTPacket(b []byte) ([]byte, int, bool) {
if len(b) < 2 {
return nil, 0, false
}
rem, mult := 0, 1
for i := 1; i < len(b) && i <= 4; i++ {
rem += int(b[i]&127) * mult
if b[i]&128 == 0 {
total := 1 + i + rem
if total < 0 || len(b) < total {
return nil, 0, false
}
pkt := make([]byte, total)
copy(pkt, b[:total])
return pkt, total, true
}
mult *= 128
}
return nil, 0, false
}
func readWSFrame(r *bufio.Reader) (payload []byte, opcode byte, err error) {
b0, err := r.ReadByte()
if err != nil {
+185
View File
@@ -0,0 +1,185 @@
package harness
import (
"bufio"
"bytes"
"encoding/binary"
"io"
"net"
"sync"
"testing"
"time"
)
var pingReq = []byte{0xC0, 0x00}
func TestSplitMQTTPacket(t *testing.T) {
t.Parallel()
if _, _, ok := splitMQTTPacket(nil); ok {
t.Fatal("empty")
}
if _, _, ok := splitMQTTPacket([]byte{0xC0}); ok {
t.Fatal("incomplete header")
}
three := concat(pingReq, pingReq, []byte{0xD0, 0x00})
pkt, n, ok := splitMQTTPacket(three)
if !ok || n != 2 || !bytes.Equal(pkt, pingReq) {
t.Fatalf("first pkt=%x n=%d ok=%v", pkt, n, ok)
}
pkt, n, ok = splitMQTTPacket(three[n:])
if !ok || n != 2 || !bytes.Equal(pkt, pingReq) {
t.Fatalf("second pkt=%x n=%d ok=%v", pkt, n, ok)
}
pkt, n, ok = splitMQTTPacket(three[4:])
if !ok || n != 2 || pkt[0] != 0xD0 {
t.Fatalf("third pkt=%x n=%d ok=%v", pkt, n, ok)
}
}
func TestWSMQTTRecvSplitsCoalescedPackets(t *testing.T) {
t.Parallel()
c, peer := pipeWS(t)
three := concat(pingReq, pingReq, []byte{0xD0, 0x00})
errCh := make(chan error, 1)
go func() { errCh <- writeWSServerFrame(peer, 0x2, three, true) }()
got := make([][]byte, 0, 3)
for i := 0; i < 3; i++ {
pkt, err := c.Recv()
if err != nil {
t.Fatalf("recv %d: %v", i, err)
}
got = append(got, pkt)
}
if err := <-errCh; err != nil {
t.Fatalf("write: %v", err)
}
want := [][]byte{pingReq, pingReq, {0xD0, 0x00}}
for i := range want {
if !bytes.Equal(got[i], want[i]) {
t.Fatalf("pkt %d = %x want %x", i, got[i], want[i])
}
}
}
func TestWSMQTTRecvContinuation(t *testing.T) {
t.Parallel()
c, peer := pipeWS(t)
errCh := make(chan error, 1)
go func() {
if err := writeWSServerFrame(peer, 0x2, []byte{0xC0}, false); err != nil {
errCh <- err
return
}
errCh <- writeWSServerFrame(peer, 0x0, []byte{0x00}, true)
}()
pkt, err := c.Recv()
if err != nil {
t.Fatalf("recv: %v", err)
}
if err := <-errCh; err != nil {
t.Fatalf("write: %v", err)
}
if !bytes.Equal(pkt, pingReq) {
t.Fatalf("pkt=%x", pkt)
}
}
func TestWSMQTTSendConcurrentLocked(t *testing.T) {
t.Parallel()
c, peer := pipeWS(t)
const nSenders = 2
const perSender = 1000
want := nSenders * perSender
var wg sync.WaitGroup
errCh := make(chan error, nSenders)
for i := 0; i < nSenders; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for j := 0; j < perSender; j++ {
if err := c.Send(pingReq); err != nil {
errCh <- err
return
}
}
}()
}
r := bufio.NewReader(peer)
_ = peer.SetDeadline(time.Now().Add(15 * time.Second))
var buf []byte
got := 0
for got < want {
payload, opcode, err := readWSFrame(r)
if err != nil {
t.Fatalf("read ws after %d pkts: %v", got, err)
}
if opcode != 0x0 && opcode != 0x2 {
continue
}
buf = append(buf, payload...)
for {
pkt, n, ok := splitMQTTPacket(buf)
if !ok {
break
}
buf = buf[n:]
if !bytes.Equal(pkt, pingReq) {
t.Fatalf("decoded %x after %d", pkt, got)
}
got++
}
}
wg.Wait()
select {
case err := <-errCh:
t.Fatalf("send: %v", err)
default:
}
if got != want {
t.Fatalf("got=%d want=%d leftover=%d", got, want, len(buf))
}
}
func pipeWS(t *testing.T) (*wsMQTT, net.Conn) {
t.Helper()
a, b := net.Pipe()
t.Cleanup(func() {
_ = a.Close()
_ = b.Close()
})
_ = a.SetDeadline(time.Now().Add(15 * time.Second))
_ = b.SetDeadline(time.Now().Add(15 * time.Second))
return &wsMQTT{conn: a, r: bufio.NewReader(a)}, b
}
func concat(parts ...[]byte) []byte {
return bytes.Join(parts, nil)
}
func writeWSServerFrame(w io.Writer, opcode byte, payload []byte, fin bool) error {
b0 := opcode & 0x0f
if fin {
b0 |= 0x80
}
header := []byte{b0}
n := len(payload)
switch {
case n < 126:
header = append(header, byte(n))
case n <= 65535:
header = append(header, 126, byte(n>>8), byte(n))
default:
var ext [8]byte
binary.BigEndian.PutUint64(ext[:], uint64(n))
header = append(header, 127)
header = append(header, ext[:]...)
}
frame := make([]byte, len(header)+n)
copy(frame, header)
copy(frame[len(header):], payload)
_, err := w.Write(frame)
return err
}