package broker import ( "bytes" "context" "errors" "io" "net" "strconv" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" "github.com/mochi-mqtt/server/v2/packets" ) func TestPublishDownBackpressureWhenQueueFull(t *testing.T) { // 无缓冲 pipe:发送 goroutine 在客户端不读时堵住,队列才能填满。 b, err := New(Options{Authenticator: AllowAuthenticator{}}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() r, w := net.Pipe() done := make(chan struct{}) go func() { defer close(done) _ = b.AttachTCP(r) }() defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() connectAndSubscribe(t, w, "ep-bp", 0) waitSession(t, b, "ep-bp") payload := bytes.Repeat([]byte("q"), 1024) var sawBP bool start := time.Now() // mochi outbound 缓冲 1024,发送 goroutine 要先填满它才会堵住,随后才轮到本地下行队列。 for i := 0; i < 1024+downQueueMax+16; i++ { ctx, cancel := context.WithTimeout(context.Background(), time.Second) err := b.PublishDown(ctx, "ep-bp", "", payload, port.PublishOpts{QoS: 1}) cancel() if errors.Is(err, ErrBackpressure) { if time.Since(start) > 50*time.Millisecond && i == 0 { t.Fatalf("first backpressure took %s", time.Since(start)) } sawBP = true break } if err != nil { t.Fatalf("publish %d: %v", i, err) } } if !sawBP { t.Fatal("expected ErrBackpressure") } } func TestSlowClientDoesNotBlockOtherPublishDown(t *testing.T) { b, err := New(Options{Authenticator: AllowAuthenticator{}}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() slowW, slowDone := acceptAndDial(t, b) fastW, fastDone := acceptAndDial(t, b) defer func() { _ = slowW.Close() _ = fastW.Close() select { case <-slowDone: case <-time.After(3 * time.Second): } select { case <-fastDone: case <-time.After(3 * time.Second): } }() writeConnect(t, slowW, "ep-slow", 30, 0) readExactPacket(t, slowW, packets.Connack, 3*time.Second) writeSubscribe(t, slowW, downTopic("ep-slow")) readExactPacket(t, slowW, packets.Suback, 3*time.Second) writeConnect(t, fastW, "ep-fast", 30, 0) readExactPacket(t, fastW, packets.Connack, 3*time.Second) writeSubscribe(t, fastW, downTopic("ep-fast")) readExactPacket(t, fastW, packets.Suback, 3*time.Second) waitSession(t, b, "ep-slow") waitSession(t, b, "ep-fast") big := bytes.Repeat([]byte("s"), 64*1024) for i := 0; i < 8; i++ { _ = b.PublishDown(context.Background(), "ep-slow", "", big, port.PublishOpts{QoS: 1}) } small := []byte(`{"v":1,"type":"resp"}`) start := time.Now() if err := b.PublishDown(context.Background(), "ep-fast", "", small, port.PublishOpts{QoS: 1}); err != nil { t.Fatalf("fast publish: %v", err) } if time.Since(start) > 100*time.Millisecond { t.Fatalf("fast PublishDown took %s", time.Since(start)) } got := readDownPayload(t, fastW, 2*time.Second) if !bytes.Equal(got, small) { t.Fatalf("fast got %q", got) } } func TestDownlinkFIFOOrder(t *testing.T) { b, w, done := startTCPClient(t, "ep-ord") defer func() { _ = b.Close() }() defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-ord", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) writeSubscribe(t, w, downTopic("ep-ord")) readExactPacket(t, w, packets.Suback, 3*time.Second) waitSession(t, b, "ep-ord") const n = 64 for i := 0; i < n; i++ { p := []byte(strconv.Itoa(i)) if err := b.PublishDown(context.Background(), "ep-ord", "", p, port.PublishOpts{QoS: 1}); err != nil { t.Fatal(err) } } for i := 0; i < n; i++ { got := readDownPayload(t, w, 3*time.Second) want := []byte(strconv.Itoa(i)) if !bytes.Equal(got, want) { t.Fatalf("order %d: got %s want %s", i, got, want) } } } func acceptAndDial(t *testing.T, b *Broker) (net.Conn, chan struct{}) { t.Helper() 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 w, done } func readDownPayload(t *testing.T, conn net.Conn, timeout time.Duration) []byte { t.Helper() deadline := time.Now().Add(timeout) for time.Now().Before(deadline) { _ = conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond)) hdr := make([]byte, 1) if _, err := io.ReadFull(conn, hdr); err != nil { continue } rem, err := readRemainingLengthConn(conn) if err != nil { continue } body := make([]byte, rem) if _, err := io.ReadFull(conn, body); err != nil { continue } if hdr[0]>>4 != packets.Publish { continue } qos := (hdr[0] >> 1) & 0x3 pk := new(packets.Packet) pk.ProtocolVersion = 5 pk.FixedHeader = packets.FixedHeader{Type: packets.Publish, Remaining: rem, Qos: qos} if err := pk.PublishDecode(body); err != nil { t.Fatal(err) } if qos > 0 { ack := packets.Packet{ FixedHeader: packets.FixedHeader{Type: packets.Puback}, ProtocolVersion: 5, PacketID: pk.PacketID, } var ab bytes.Buffer _ = ack.PubackEncode(&ab) _, _ = conn.Write(ab.Bytes()) } return pk.Payload } t.Fatal("timeout waiting publish") return nil }