package broker import ( "bytes" "context" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" "github.com/mochi-mqtt/server/v2/packets" ) type downItem struct { payload []byte qos byte disconnect port.DisconnectReason // 非空表示该帧写出后断开(B-04) sent chan struct{} } func (st *connState) startDownLoop(b *Broker) { st.mu.Lock() if st.downCh != nil { st.mu.Unlock() return } st.downCh = make(chan downItem, downQueueMax) st.downStop = make(chan struct{}) st.downDone = make(chan struct{}) st.mu.Unlock() go st.downLoop(b) } func (st *connState) stopDownLoop() { st.mu.Lock() stop := st.downStop done := st.downDone ch := st.downCh st.mu.Unlock() if stop == nil { return } select { case <-stop: default: close(stop) } if done != nil { select { case <-done: case <-time.After(2 * time.Second): } } if ch != nil { for { select { case item := <-ch: st.downBytes.Add(-int64(len(item.payload))) default: return } } } } func (st *connState) enqueueDown(item downItem) error { st.mu.Lock() ch := st.downCh stop := st.downStop closing := st.closing st.mu.Unlock() if ch == nil || stop == nil { return ErrNoConnection } select { case <-stop: return ErrNoConnection default: } if closing && item.disconnect == "" { return ErrNoConnection } n := int64(len(item.payload)) for { cur := st.downBytes.Load() if cur+n > downQueueBytes { return ErrBackpressure } if st.downBytes.CompareAndSwap(cur, cur+n) { break } } select { case ch <- item: return nil default: st.downBytes.Add(-n) return ErrBackpressure } } func (st *connState) downLoop(b *Broker) { defer close(st.downDone) for { select { case <-st.downStop: return case item, ok := <-st.downCh: if !ok { return } st.downBytes.Add(-int64(len(item.payload))) st.sendOne(b, item) } } } func (st *connState) sendOne(b *Broker, item downItem) { if b.closed.Load() { st.signalSent(item) return } large := len(item.payload) > largeFrameBytes if large { if err := b.acquireLarge(context.Background()); err != nil { if b.onDrop != nil { b.onDrop(context.Background(), st.endpointID, st.connID, item.payload) } st.signalSent(item) return } st.mu.Lock() st.largePending++ st.mu.Unlock() } topic := downTopic(st.endpointID) var waitCh chan struct{} if item.disconnect != "" { waitCh = st.armWriteWait(item.payload) } st.wirePending.Add(1) var err error for !b.closed.Load() { select { case <-st.downStop: err = ErrNoConnection default: err = b.server.Publish(topic, item.payload, false, item.qos) if err == nil { break } select { case <-st.downStop: err = ErrNoConnection case <-time.After(2 * time.Millisecond): continue } } 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) } if err != nil && b.onDrop != nil { b.onDrop(context.Background(), st.endpointID, st.connID, item.payload) } st.signalSent(item) if err == nil && item.disconnect != "" { st.waitWriteDone(waitCh) _ = b.Disconnect(context.Background(), st.endpointID, st.connID, item.disconnect) } } 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) { if item.sent == nil { return } select { case <-item.sent: default: close(item.sent) } }