package broker import ( "bytes" "context" "errors" "io" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" "github.com/mochi-mqtt/server/v2/packets" ) func TestPublishDownRequiresDownSubscription(t *testing.T) { b, w, done := startTCPClient(t, "ep-nosub") defer func() { _ = b.Close() }() defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-nosub", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) waitSession(t, b, "ep-nosub") err := b.PublishDown(context.Background(), "ep-nosub", "", []byte(`{"v":1}`), port.PublishOpts{QoS: 0}) if !errors.Is(err, ErrNotSubscribed) { t.Fatalf("err=%v want ErrNotSubscribed", err) } } func TestPublishDownRejectsStaleConnID(t *testing.T) { b, w, done := startTCPClient(t, "ep-stale") defer func() { _ = b.Close() }() defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-stale", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) writeSubscribe(t, w, downTopic("ep-stale")) readExactPacket(t, w, packets.Suback, 3*time.Second) waitSession(t, b, "ep-stale") err := b.PublishDown(context.Background(), "ep-stale", "dead-conn", []byte(`{"v":1}`), port.PublishOpts{QoS: 0}) if !errors.Is(err, ErrNoConnection) { t.Fatalf("err=%v want ErrNoConnection", err) } } func TestPublishThenDisconnectWritesThenCloses(t *testing.T) { b, w, done := startTCPClient(t, "ep-ptd") defer func() { _ = b.Close() }() defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-ptd", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) writeSubscribe(t, w, downTopic("ep-ptd")) readExactPacket(t, w, packets.Suback, 3*time.Second) waitSession(t, b, "ep-ptd") payload := []byte(`{"v":1,"type":"fatal","reason":"disabled"}`) if err := b.PublishThenDisconnect(context.Background(), "ep-ptd", "", payload, 1, port.DisconnectFatal); err != nil { t.Fatal(err) } got := readDownPayload(t, w, 3*time.Second) if !bytes.Equal(got, payload) { t.Fatalf("got %s", got) } _ = w.SetReadDeadline(time.Now().Add(3 * time.Second)) buf := make([]byte, 64) n, err := io.ReadAtLeast(w, buf, 2) if err != nil && n == 0 { return // 连接已关 } if n > 0 && buf[0]>>4 == packets.Disconnect { return } } func TestPublishThenDisconnectAfterQueuedFrame(t *testing.T) { b, w, done := startTCPClient(t, "ep-ptd-q") defer func() { _ = b.Close() }() defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-ptd-q", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) writeSubscribe(t, w, downTopic("ep-ptd-q")) readExactPacket(t, w, packets.Suback, 3*time.Second) waitSession(t, b, "ep-ptd-q") first := []byte(`{"v":1,"type":"msg","id":"queued-ahead"}`) fatal := []byte(`{"v":1,"type":"fatal","reason":"disabled"}`) if err := b.PublishDown(context.Background(), "ep-ptd-q", "", first, port.PublishOpts{QoS: 1}); err != nil { t.Fatal(err) } if err := b.PublishThenDisconnect(context.Background(), "ep-ptd-q", "", fatal, 1, port.DisconnectFatal); err != nil { t.Fatal(err) } gotFirst := readDownPayload(t, w, 3*time.Second) if !bytes.Equal(gotFirst, first) { t.Fatalf("first got %s", gotFirst) } gotFatal := readDownPayload(t, w, 3*time.Second) if !bytes.Equal(gotFatal, fatal) { t.Fatalf("fatal got %s want %s (disconnected before fatal frame)", gotFatal, fatal) } _ = w.SetReadDeadline(time.Now().Add(3 * time.Second)) buf := make([]byte, 64) n, err := io.ReadAtLeast(w, buf, 2) if err != nil && n == 0 { return } if n > 0 && buf[0]>>4 == packets.Disconnect { return } } func TestShutdownUsesServerShuttingDown(t *testing.T) { b, w, done := startTCPClient(t, "ep-shut") defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-shut", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) waitSession(t, b, "ep-shut") ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() if err := b.Shutdown(ctx); err != nil { t.Fatal(err) } _ = w.SetReadDeadline(time.Now().Add(2 * time.Second)) buf := make([]byte, 32) n, err := io.ReadAtLeast(w, buf, 2) if err != nil && n == 0 { return } if n > 0 && buf[0]>>4 != packets.Disconnect { t.Fatalf("want disconnect got %x", buf[:n]) } } func TestShutdownWithBacklogDeliversServerShuttingDown(t *testing.T) { b, w, done := startTCPClient(t, "ep-shut-bl") defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-shut-bl", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) writeSubscribe(t, w, downTopic("ep-shut-bl")) readExactPacket(t, w, packets.Suback, 3*time.Second) waitSession(t, b, "ep-shut-bl") payload := bytes.Repeat([]byte("b"), 1024) for i := 0; i < 8; i++ { if err := b.PublishDown(context.Background(), "ep-shut-bl", "", payload, port.PublishOpts{QoS: 1}); err != nil { t.Fatalf("publish %d: %v", i, err) } } saw8B := make(chan bool, 1) go func() { deadline := time.Now().Add(3 * time.Second) for time.Now().Before(deadline) { _ = 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 } switch hdr[0] >> 4 { case packets.Publish: qos := (hdr[0] >> 1) & 0x3 if qos > 0 { 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 { ack := packets.Packet{ FixedHeader: packets.FixedHeader{Type: packets.Puback}, ProtocolVersion: 5, PacketID: pk.PacketID, } var ab bytes.Buffer _ = ack.PubackEncode(&ab) _, _ = w.Write(ab.Bytes()) } } case packets.Disconnect: if rem >= 1 && body[0] == packets.ErrServerShuttingDown.Code { saw8B <- true return } } } saw8B <- false }() time.Sleep(20 * time.Millisecond) ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) defer cancel() err := b.Shutdown(ctx) got := false select { case got = <-saw8B: case <-time.After(4 * time.Second): t.Fatal("reader hung") } if got { if err != nil && !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("shutdown after 0x8B: %v", err) } return } if err == nil { t.Fatal("expected 0x8B or shutdown deadline error, got neither") } if !errors.Is(err, context.DeadlineExceeded) && !errors.Is(err, context.Canceled) { t.Fatalf("shutdown err=%v want deadline/cancel when 0x8B not seen", err) } } func TestShutdownCancelledContextReturnsQuickly(t *testing.T) { b, w, done := startTCPClient(t, "ep-shut-cancel") defer func() { _ = w.Close() select { case <-done: case <-time.After(3 * time.Second): } }() writeConnect(t, w, "ep-shut-cancel", 30, 0) readExactPacket(t, w, packets.Connack, 3*time.Second) waitSession(t, b, "ep-shut-cancel") ctx, cancel := context.WithCancel(context.Background()) cancel() start := time.Now() _ = b.Shutdown(ctx) if time.Since(start) > 500*time.Millisecond { t.Fatalf("cancelled shutdown took %s", time.Since(start)) } } func TestWaitConnsQuietSeesInSend(t *testing.T) { b, err := New(Options{}) if err != nil { t.Fatal(err) } defer func() { _ = b.Close() }() st := &connState{downCh: make(chan downItem, 1)} st.inSend.Store(1) ctx, cancel := context.WithTimeout(context.Background(), 40*time.Millisecond) defer cancel() if b.waitConnsQuiet(ctx, []*connState{st}) { t.Fatal("inSend should keep shutdown from treating the conn as quiet") } st.inSend.Store(0) ctx2, cancel2 := context.WithTimeout(context.Background(), 200*time.Millisecond) defer cancel2() if !b.waitConnsQuiet(ctx2, []*connState{st}) { t.Fatal("quiet when inSend is 0 and queues are empty") } } func TestEffectivePayloadLimitSubtractsOverhead(t *testing.T) { got := EffectivePayloadLimit(200, 0) if got != 200-packetOverheadBudget { t.Fatalf("got %d", got) } got = EffectivePayloadLimit(200, 50) if got != 50 { t.Fatalf("got %d want 50", got) } }