226 lines
4.7 KiB
Go
226 lines
4.7 KiB
Go
package broker
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"io"
|
|
"net"
|
|
"runtime"
|
|
"sync"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
|
"github.com/mochi-mqtt/server/v2/packets"
|
|
)
|
|
|
|
func TestReceiveMaximumDoesNotDeadlockPublish(t *testing.T) {
|
|
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = b.Close() }()
|
|
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = ln.Close() }()
|
|
|
|
clientDone := make(chan struct{})
|
|
go func() {
|
|
defer close(clientDone)
|
|
c, accErr := ln.Accept()
|
|
if accErr != nil {
|
|
return
|
|
}
|
|
_ = b.AttachTCP(c)
|
|
}()
|
|
|
|
w, err := net.Dial("tcp", ln.Addr().String())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = w.Close() }()
|
|
|
|
endpoint := "ep-rm-quota"
|
|
writeConnectFull(t, w, endpoint, 30, 0, 20)
|
|
readExactPacket(t, w, packets.Connack, 3*time.Second)
|
|
writeSubscribe(t, w, downTopic(endpoint))
|
|
readExactPacket(t, w, packets.Suback, 3*time.Second)
|
|
|
|
deadline := time.Now().Add(3 * time.Second)
|
|
for {
|
|
if _, ok := b.ConnInfoOf(endpoint); ok {
|
|
break
|
|
}
|
|
if time.Now().After(deadline) {
|
|
t.Fatal("session not established")
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
|
|
var received atomic.Int64
|
|
var writeMu sync.Mutex
|
|
stop := make(chan struct{})
|
|
var stopOnce sync.Once
|
|
halt := func() { stopOnce.Do(func() { close(stop) }) }
|
|
defer halt()
|
|
|
|
go func() {
|
|
for {
|
|
select {
|
|
case <-stop:
|
|
return
|
|
default:
|
|
}
|
|
_ = w.SetReadDeadline(time.Now().Add(200 * time.Millisecond))
|
|
hdr := make([]byte, 1)
|
|
if _, err := io.ReadFull(w, hdr); err != nil {
|
|
continue
|
|
}
|
|
rem, err := readRemainingLengthConn(w)
|
|
if err != nil {
|
|
continue
|
|
}
|
|
body := make([]byte, rem)
|
|
if _, err := io.ReadFull(w, body); err != nil {
|
|
continue
|
|
}
|
|
typ := hdr[0] >> 4
|
|
if typ != packets.Publish {
|
|
continue
|
|
}
|
|
qos := (hdr[0] >> 1) & 0x3
|
|
received.Add(1)
|
|
if qos == 0 {
|
|
continue
|
|
}
|
|
pk := new(packets.Packet)
|
|
pk.ProtocolVersion = 5
|
|
pk.FixedHeader = packets.FixedHeader{Type: packets.Publish, Remaining: rem, Qos: qos}
|
|
if decErr := pk.PublishDecode(body); decErr != nil {
|
|
continue
|
|
}
|
|
ack := packets.Packet{
|
|
FixedHeader: packets.FixedHeader{Type: packets.Puback},
|
|
ProtocolVersion: 5,
|
|
PacketID: pk.PacketID,
|
|
}
|
|
var ab bytes.Buffer
|
|
_ = ack.PubackEncode(&ab)
|
|
writeMu.Lock()
|
|
_, _ = w.Write(ab.Bytes())
|
|
writeMu.Unlock()
|
|
}
|
|
}()
|
|
|
|
go func() {
|
|
tick := time.NewTicker(2 * time.Millisecond)
|
|
defer tick.Stop()
|
|
pk := packets.Packet{
|
|
FixedHeader: packets.FixedHeader{Type: packets.Pingreq},
|
|
ProtocolVersion: 5,
|
|
}
|
|
var buf bytes.Buffer
|
|
_ = pk.PingreqEncode(&buf)
|
|
ping := append([]byte(nil), buf.Bytes()...)
|
|
for {
|
|
select {
|
|
case <-stop:
|
|
return
|
|
case <-tick.C:
|
|
writeMu.Lock()
|
|
_, _ = w.Write(ping)
|
|
writeMu.Unlock()
|
|
}
|
|
}
|
|
}()
|
|
|
|
payload := []byte(`{"v":1,"type":"resp","rid":"x"}`)
|
|
pubCtx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
var wg sync.WaitGroup
|
|
for i := 0; i < 8; i++ {
|
|
wg.Add(1)
|
|
go func() {
|
|
defer wg.Done()
|
|
for pubCtx.Err() == nil {
|
|
_ = b.PublishDown(pubCtx, endpoint, "", payload, port.PublishOpts{QoS: 1})
|
|
}
|
|
}()
|
|
}
|
|
|
|
runFor := 5 * time.Second
|
|
watch := 2 * time.Second
|
|
start := time.Now()
|
|
last := received.Load()
|
|
lastChange := time.Now()
|
|
for time.Since(start) < runFor {
|
|
time.Sleep(50 * time.Millisecond)
|
|
n := received.Load()
|
|
if n > last {
|
|
last = n
|
|
lastChange = time.Now()
|
|
}
|
|
if time.Since(lastChange) > watch {
|
|
buf := make([]byte, 1<<20)
|
|
nstack := runtime.Stack(buf, true)
|
|
halt()
|
|
cancel()
|
|
_ = w.Close()
|
|
t.Fatalf("progress stalled at %d after %s\n%s", last, time.Since(lastChange), buf[:nstack])
|
|
}
|
|
}
|
|
cancel()
|
|
wg.Wait()
|
|
halt()
|
|
if last < 100 {
|
|
t.Fatalf("too few publishes delivered: %d", last)
|
|
}
|
|
|
|
_ = w.Close()
|
|
select {
|
|
case <-clientDone:
|
|
case <-time.After(3 * time.Second):
|
|
}
|
|
}
|
|
|
|
func readExactPacket(t *testing.T, conn net.Conn, wantType byte, timeout time.Duration) {
|
|
t.Helper()
|
|
_ = conn.SetReadDeadline(time.Now().Add(timeout))
|
|
hdr := make([]byte, 1)
|
|
if _, err := io.ReadFull(conn, hdr); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if hdr[0]>>4 != wantType {
|
|
t.Fatalf("want packet type %d got %d", wantType, hdr[0]>>4)
|
|
}
|
|
rem, err := readRemainingLengthConn(conn)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
body := make([]byte, rem)
|
|
if _, err := io.ReadFull(conn, body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func readRemainingLengthConn(r io.Reader) (int, error) {
|
|
var mul uint32 = 1
|
|
var value uint32
|
|
for i := 0; i < 4; i++ {
|
|
var b [1]byte
|
|
if _, err := io.ReadFull(r, b[:]); err != nil {
|
|
return 0, err
|
|
}
|
|
value += uint32(b[0]&127) * mul
|
|
if b[0]&128 == 0 {
|
|
return int(value), nil
|
|
}
|
|
mul *= 128
|
|
}
|
|
return 0, io.ErrUnexpectedEOF
|
|
}
|