295 lines
5.4 KiB
Go
295 lines
5.4 KiB
Go
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
|
|
}
|
|
st.inSend.Add(1)
|
|
defer st.inSend.Add(-1)
|
|
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)
|
|
}
|
|
}
|