package broker import ( "bytes" "context" "io" "net" "sync" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" "github.com/mochi-mqtt/server/v2/packets" ) func TestLargeFrameQuotaReleasedOnPuback(t *testing.T) { b, w, done := startTCPClient(t, "ep-large-ack") defer func() { _ = b.Close() }() defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-large-ack", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) writeSubscribe(t, w, downTopic("ep-large-ack")) readExactPacket(t, w, packets.Suback, 3*time.Second) waitSession(t, b, "ep-large-ack") stop := make(chan struct{}) defer close(stop) var writeMu sync.Mutex go autoPuback(w, stop, &writeMu) payload := bytes.Repeat([]byte("x"), 70*1024) for i := 0; i < 65; i++ { ctx, cancel := context.WithTimeout(context.Background(), time.Second) err := b.PublishDown(ctx, "ep-large-ack", "", payload, port.PublishOpts{QoS: 1}) cancel() if err != nil { t.Fatalf("publish %d: %v", i+1, err) } } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { if len(b.largeSem) == 0 { return } time.Sleep(10 * time.Millisecond) } t.Fatalf("slots still held: %d", len(b.largeSem)) } func TestLargeFrameQuotaReleasedOnDisconnect(t *testing.T) { b, w, done := startTCPClient(t, "ep-large-disc") defer func() { _ = b.Close() }() writeConnect(t, w, "ep-large-disc", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) writeSubscribe(t, w, downTopic("ep-large-disc")) readExactPacket(t, w, packets.Suback, 3*time.Second) waitSession(t, b, "ep-large-disc") go func() { buf := make([]byte, 32*1024) for { _, err := w.Read(buf) if err != nil { return } } }() payload := bytes.Repeat([]byte("y"), 70*1024) ctx, cancel := context.WithTimeout(context.Background(), time.Second) if err := b.PublishDown(ctx, "ep-large-disc", "", payload, port.PublishOpts{QoS: 1}); err != nil { cancel() t.Fatal(err) } cancel() if len(b.largeSem) == 0 { t.Fatal("expected a held large slot before disconnect") } _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): t.Fatal("client attach did not return") } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { if len(b.largeSem) == 0 { return } time.Sleep(10 * time.Millisecond) } t.Fatalf("slots after disconnect: %d", len(b.largeSem)) } func TestLargeFrameQuotaReleasedWithoutSubscriber(t *testing.T) { b, w, done := startTCPClient(t, "ep-large-nosub") defer func() { _ = b.Close() }() defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-large-nosub", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) waitSession(t, b, "ep-large-nosub") payload := bytes.Repeat([]byte("z"), 70*1024) ctx, cancel := context.WithTimeout(context.Background(), time.Second) err := b.PublishDown(ctx, "ep-large-nosub", "", payload, port.PublishOpts{QoS: 1}) cancel() if err != nil { t.Fatal(err) } if n := len(b.largeSem); n != 0 { t.Fatalf("held slots without subscriber: %d", n) } } func startTCPClient(t *testing.T, _ string) (*Broker, net.Conn, chan struct{}) { t.Helper() b, err := New(Options{Authenticator: AllowAuthenticator{}}) if err != nil { t.Fatal(err) } ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = ln.Close() }) done := make(chan struct{}) go func() { defer close(done) 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) } return b, w, done } func waitSession(t *testing.T, b *Broker, endpoint string) { t.Helper() deadline := time.Now().Add(3 * time.Second) for time.Now().Before(deadline) { if _, ok := b.ConnInfoOf(endpoint); ok { return } time.Sleep(5 * time.Millisecond) } t.Fatal("session not established") } func autoPuback(w net.Conn, stop <-chan struct{}, writeMu *sync.Mutex) { 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 } if hdr[0]>>4 != packets.Publish { continue } qos := (hdr[0] >> 1) & 0x3 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() } }