fix: 积压时等本帧写出再断开并让 Shutdown 等待 0x8B
This commit is contained in:
+115
-12
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user