fix: 积压时等本帧写出再断开并让 Shutdown 等待 0x8B

This commit is contained in:
Nixevol
2026-09-30 19:20:24 +08:00
parent 0c9b459fb1
commit a14ca0d08d
6 changed files with 375 additions and 22 deletions
+115 -12
View File
@@ -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) {