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 }