From a14ca0d08d6706158b33b9df297ccd820d658e6c Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 19:20:24 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E7=A7=AF=E5=8E=8B=E6=97=B6=E7=AD=89?= =?UTF-8?q?=E6=9C=AC=E5=B8=A7=E5=86=99=E5=87=BA=E5=86=8D=E6=96=AD=E5=BC=80?= =?UTF-8?q?=E5=B9=B6=E8=AE=A9=20Shutdown=20=E7=AD=89=E5=BE=85=200x8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- cmd/nixmsg/serve.go | 4 +- docs/DEVIATIONS.md | 10 ++ internal/broker/b04_b06_b08_test.go | 158 ++++++++++++++++++++++++++++ internal/broker/broker.go | 92 +++++++++++++++- internal/broker/downlink.go | 127 +++++++++++++++++++--- internal/broker/hooks.go | 6 +- 6 files changed, 375 insertions(+), 22 deletions(-) diff --git a/cmd/nixmsg/serve.go b/cmd/nixmsg/serve.go index ebadc52..3cb821d 100644 --- a/cmd/nixmsg/serve.go +++ b/cmd/nixmsg/serve.go @@ -335,7 +335,9 @@ func runServe(ctx context.Context, cfg config.Config) error { drainCancel() shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second) - _ = brk.Shutdown(shutCtx) + if shutErr := brk.Shutdown(shutCtx); shutErr != nil && !errors.Is(shutErr, context.DeadlineExceeded) && !errors.Is(shutErr, context.Canceled) { + slog.Error("broker shutdown", "err", shutErr) + } shutCancel() secondDrain := drainBudget - time.Since(drainStart) diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 7deb4c5..81c9727 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1718,3 +1718,13 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 - 原因:原先 12 项备注写着未穷尽仍标通过,交付说明写成「通过 23」。 - 备选方案:为每个未测子项补验收用例(本波不做,避免为变绿放松断言)。 - 影响:汇总改为通过 19、部分通过 4(F03/F08/F21/F22)、失败 0。F19 仍引用仓库内 SDK 清单、本波不重跑。 + +### 复审修复 R3-02 + +1. **积压时 fatal/logout 与停机 0x8B 须等本帧写出** + - 日期:2026-09-30 + - 原条款:Gitea #66;B-04 / B-08。 + - 实际做法:带断开的下行帧用本帧 `OnPacketSent` 完成信号(优先 packet id,否则按载荷匹配),不再用连接级 `sentPub` 总数。`Shutdown` 在 ctx 未取消时先等下行队列与 `wirePending` 排空,再 `DisconnectClient` 发 `0x8B` 并在截止前等连接拆掉;ctx 已取消则发完即 `Close`。`serve` 仍给 5 秒预算并记录非超时错误。未合 `feat/fix-3-downlink-deadlock`。 + - 原因:前面 PUBLISH 的 `OnPacketSent` 会让总数等待提前返回;`Shutdown` 对 ctx 非阻塞 select 使 5 秒预算用不上,有 outbound 积压时 `0x8B` 只进 outbuf 随 `Stop` 丢掉。 + - 备选方案:恢复固定 `Sleep`(否决);改 `PublishDown` 签名(否决)。 + - 影响:队列/outbound 有积压时 fatal、logout 先到客户端再断开;停机在预算内尽量发出 `0x8B`,超时返回 ctx 错误而非空等。 diff --git a/internal/broker/b04_b06_b08_test.go b/internal/broker/b04_b06_b08_test.go index 806f7e7..27e1424 100644 --- a/internal/broker/b04_b06_b08_test.go +++ b/internal/broker/b04_b06_b08_test.go @@ -87,6 +87,50 @@ func TestPublishThenDisconnectWritesThenCloses(t *testing.T) { } } +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() { @@ -116,6 +160,120 @@ func TestShutdownUsesServerShuttingDown(t *testing.T) { } } +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 TestEffectivePayloadLimitSubtractsOverhead(t *testing.T) { got := EffectivePayloadLimit(200, 0) if got != 200-packetOverheadBudget { diff --git a/internal/broker/broker.go b/internal/broker/broker.go index 7530eab..a42fdb7 100644 --- a/internal/broker/broker.go +++ b/internal/broker/broker.go @@ -141,9 +141,14 @@ type connState struct { downStop chan struct{} downDone chan struct{} downBytes atomic.Int64 - sentPub atomic.Int64 + wirePending atomic.Int64 // Publish 入 mochi outbound 后、OnPacketSent 前 mu sync.Mutex + // 带断开的下行帧:只等本帧 OnPacketSent,不用连接级计数。 + writeWaitCh chan struct{} + writeWaitPayload []byte + writeWaitPID uint16 // 非 0 时优先按 packet id 匹配 + handshakeTimer *time.Timer } @@ -232,28 +237,105 @@ func (b *Broker) Close() error { } // Shutdown 向所有连接发 MQTT 5 0x8B 后关闭。完整 HTTP 停机顺序见 L-03。 +// ctx 未取消时先等下行队列与 wirePending 排空,再 DisconnectClient(此时 outbound 空, +// 0x8B 直写套接字),并在截止前等连接拆掉;ctx 已取消则发完即 Close,不等待。 func (b *Broker) Shutdown(ctx context.Context) error { if b.closed.Load() { return nil } + if ctx == nil { + ctx = context.Background() + } + b.connsMu.RLock() + states := make([]*connState, 0, len(b.byClient)) clients := make([]*mqtt.Client, 0, len(b.byClient)) - for cl := range b.byClient { + for cl, st := range b.byClient { if cl != nil { clients = append(clients, cl) } + if st != nil { + states = append(states, st) + } } b.connsMu.RUnlock() + + for _, st := range states { + st.mu.Lock() + st.closing = true + st.mu.Unlock() + } + + alreadyCancelled := false + select { + case <-ctx.Done(): + alreadyCancelled = true + default: + } + + var waitErr error + if !alreadyCancelled { + if !b.waitConnsQuiet(ctx, states) { + waitErr = ctx.Err() + } + } + for _, cl := range clients { _ = b.server.DisconnectClient(cl, packets.ErrServerShuttingDown) } - if ctx != nil { + + if !alreadyCancelled && waitErr == nil { + waitErr = b.waitConnsGone(ctx) + } + + closeErr := b.Close() + if waitErr != nil { + return waitErr + } + return closeErr +} + +func (b *Broker) waitConnsQuiet(ctx context.Context, states []*connState) bool { + for { + quiet := true + for _, st := range states { + if st.wirePending.Load() > 0 { + quiet = false + break + } + st.mu.Lock() + ch := st.downCh + st.mu.Unlock() + if ch != nil && len(ch) > 0 { + quiet = false + break + } + } + if quiet { + return true + } select { case <-ctx.Done(): - default: + return false + case <-time.After(2 * time.Millisecond): + } + } +} + +func (b *Broker) waitConnsGone(ctx context.Context) error { + for { + b.connsMu.RLock() + n := len(b.byClient) + b.connsMu.RUnlock() + if n == 0 { + return nil + } + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(2 * time.Millisecond): } } - return b.Close() } // AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。 diff --git a/internal/broker/downlink.go b/internal/broker/downlink.go index 3c5d2d5..edb3ee8 100644 --- a/internal/broker/downlink.go +++ b/internal/broker/downlink.go @@ -1,10 +1,12 @@ package broker import ( + "bytes" "context" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" + "github.com/mochi-mqtt/server/v2/packets" ) type downItem struct { @@ -130,7 +132,11 @@ func (st *connState) sendOne(b *Broker, item downItem) { st.mu.Unlock() } topic := downTopic(st.endpointID) - before := st.sentPub.Load() + var waitCh chan struct{} + if item.disconnect != "" { + waitCh = st.armWriteWait(item.payload) + } + st.wirePending.Add(1) var err error for !b.closed.Load() { select { @@ -150,6 +156,14 @@ func (st *connState) sendOne(b *Broker, item downItem) { } break } + if err != nil { + st.wirePending.Add(-1) + st.clearWriteWait(waitCh) + } else if waitCh != nil { + if pid, ok := st.lookupInflightPID(item.payload); ok { + st.setWriteWaitPID(waitCh, pid) + } + } if large { b.finishLargePublish(st) } @@ -158,23 +172,112 @@ func (st *connState) sendOne(b *Broker, item downItem) { } st.signalSent(item) if err == nil && item.disconnect != "" { - st.waitPacketWritten(before) + st.waitWriteDone(waitCh) _ = b.Disconnect(context.Background(), st.endpointID, st.connID, item.disconnect) } } -func (st *connState) waitPacketWritten(before int64) { - deadline := time.Now().Add(2 * time.Second) - for time.Now().Before(deadline) { - if st.sentPub.Load() > before { - return - } - select { - case <-st.downStop: - return - case <-time.After(2 * time.Millisecond): +func (st *connState) armWriteWait(payload []byte) chan struct{} { + ch := make(chan struct{}) + st.mu.Lock() + st.writeWaitCh = ch + st.writeWaitPayload = payload + st.writeWaitPID = 0 + st.mu.Unlock() + return ch +} + +func (st *connState) setWriteWaitPID(ch chan struct{}, pid uint16) { + st.mu.Lock() + if st.writeWaitCh == ch { + st.writeWaitPID = pid + } + st.mu.Unlock() +} + +func (st *connState) clearWriteWait(ch chan struct{}) { + if ch == nil { + return + } + st.mu.Lock() + if st.writeWaitCh == ch { + st.writeWaitCh = nil + st.writeWaitPayload = nil + st.writeWaitPID = 0 + } + st.mu.Unlock() +} + +func (st *connState) waitWriteDone(ch chan struct{}) { + if ch == nil { + return + } + defer st.clearWriteWait(ch) + deadline := time.NewTimer(2 * time.Second) + defer deadline.Stop() + select { + case <-ch: + case <-st.downStop: + case <-deadline.C: + } +} + +func (st *connState) notePacketSent(pk packets.Packet) { + if pk.FixedHeader.Type == packets.Publish { + for { + cur := st.wirePending.Load() + if cur <= 0 { + break + } + if st.wirePending.CompareAndSwap(cur, cur-1) { + break + } } } + st.mu.Lock() + ch := st.writeWaitCh + pid := st.writeWaitPID + want := st.writeWaitPayload + st.mu.Unlock() + if ch == nil || pk.FixedHeader.Type != packets.Publish { + return + } + match := false + if pid != 0 { + match = pk.PacketID == pid + } else if want != nil { + match = bytes.Equal(pk.Payload, want) + } + if !match { + return + } + st.mu.Lock() + if st.writeWaitCh == ch { + st.writeWaitCh = nil + st.writeWaitPayload = nil + st.writeWaitPID = 0 + } + st.mu.Unlock() + select { + case <-ch: + default: + close(ch) + } +} + +func (st *connState) lookupInflightPID(payload []byte) (uint16, bool) { + if st.client == nil || st.client.State.Inflight == nil { + return 0, false + } + for _, pk := range st.client.State.Inflight.GetAll(false) { + if pk.FixedHeader.Type != packets.Publish { + continue + } + if bytes.Equal(pk.Payload, payload) { + return pk.PacketID, pk.PacketID != 0 + } + } + return 0, false } func (st *connState) signalSent(item downItem) { diff --git a/internal/broker/hooks.go b/internal/broker/hooks.go index 567a31b..674f6d5 100644 --- a/internal/broker/hooks.go +++ b/internal/broker/hooks.go @@ -175,14 +175,11 @@ func (h *nixHook) OnSubscribed(cl *mqtt.Client, pk packets.Packet, reasonCodes [ } func (h *nixHook) OnPacketSent(cl *mqtt.Client, pk packets.Packet, _ []byte) { - if pk.FixedHeader.Type != packets.Publish { - return - } h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st != nil { - st.sentPub.Add(1) + st.notePacketSent(pk) } } @@ -254,6 +251,7 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) { lk.Lock() st.stopDownLoop() lk.Unlock() + st.wirePending.Store(0) h.b.releaseAllLarge(st) h.b.cancelHandshakeDeadline(st.endpointID, st.connID)