From da43385c4f8cec9a0da795d7f4a6a803411c4aa2 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 14:58:45 +0800 Subject: [PATCH 1/4] =?UTF-8?q?fix:=20=E5=9C=A8=20OnConnect=20=E6=8A=8A=20?= =?UTF-8?q?mochi=20=E5=8F=91=E9=80=81=E9=85=8D=E9=A2=9D=E7=BD=AE=200=20?= =?UTF-8?q?=E8=A7=84=E9=81=BF=E6=AD=BB=E9=94=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 9 ++ internal/broker/b01_quota_test.go | 225 ++++++++++++++++++++++++++++++ internal/broker/broker_test.go | 6 + internal/broker/hooks.go | 10 ++ 4 files changed, 250 insertions(+) create mode 100644 internal/broker/b01_quota_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index ac2fa7e..2f4f30c 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1171,3 +1171,12 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 5. **建议的正确方向** - 在 broker 把对本连接的下行 `InjectPacket` 与上行 worker 解耦:上行读循环先写完 PUBACK,处理 `HandleUplink` 期间不要同步向本连接注入;handler 返回后再发 `resp` 和 `group_event`。不要靠固定 `Sleep`。`InlineClient: true` 保持,`OnPublish` 对 InlineClient 继续放行。 - 覆盖 presence 等其他同步 `PublishDown`,而不只包一层 `emit`。 + +### 复审修复 B-01 + +- 日期:2026-09-30 +- 原条款:DEVELOPMENT 第 5 节装配 mochi;未写客户端 Receive Maximum。Gitea #8。 +- 实际做法:`OnConnect` 在心跳校正后调用 `cl.State.Inflight.ResetSendQuota(0)`,不 fork mochi。CONNECT 声明的 Receive Maximum 小于 256 时打 warn,连接仍接受。应用层窗口(推送 32、回执 64、在途 resp 等)约束未确认的 QoS 1。 +- 原因:mochi v2.7.9 在 `sendQuota>0` 时走 `NextImmediate` 递归读锁,并可因补发后删除 inflight 泄漏配额;已验证置 0 绕开整条路径。 +- 备选方案:fork 修补 mochi(只修递归读锁仍观察到停滞)。 +- 影响:服务端不再执行客户端 Receive Maximum;裸设备若带过小的 Receive Maximum,实际在途可能超过该值。 diff --git a/internal/broker/b01_quota_test.go b/internal/broker/b01_quota_test.go new file mode 100644 index 0000000..3c1f0f4 --- /dev/null +++ b/internal/broker/b01_quota_test.go @@ -0,0 +1,225 @@ +package broker + +import ( + "bytes" + "context" + "io" + "net" + "runtime" + "sync" + "sync/atomic" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" + "github.com/mochi-mqtt/server/v2/packets" +) + +func TestReceiveMaximumDoesNotDeadlockPublish(t *testing.T) { + b, err := New(Options{Authenticator: AllowAuthenticator{}}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b.Close() }() + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer func() { _ = ln.Close() }() + + clientDone := make(chan struct{}) + go func() { + defer close(clientDone) + c, accErr := ln.Accept() + if accErr != nil { + return + } + _ = b.AttachTCP(c) + }() + + w, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + t.Fatal(err) + } + defer func() { _ = w.Close() }() + + endpoint := "ep-rm-quota" + writeConnectFull(t, w, endpoint, 30, 0, 20) + readExactPacket(t, w, packets.Connack, 3*time.Second) + writeSubscribe(t, w, downTopic(endpoint)) + readExactPacket(t, w, packets.Suback, 3*time.Second) + + deadline := time.Now().Add(3 * time.Second) + for { + if _, ok := b.ConnInfoOf(endpoint); ok { + break + } + if time.Now().After(deadline) { + t.Fatal("session not established") + } + time.Sleep(5 * time.Millisecond) + } + + var received atomic.Int64 + var writeMu sync.Mutex + stop := make(chan struct{}) + var stopOnce sync.Once + halt := func() { stopOnce.Do(func() { close(stop) }) } + defer halt() + + go func() { + for { + select { + case <-stop: + return + default: + } + _ = 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 + } + typ := hdr[0] >> 4 + if typ != packets.Publish { + continue + } + qos := (hdr[0] >> 1) & 0x3 + received.Add(1) + if qos == 0 { + continue + } + 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 { + continue + } + ack := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Puback}, + ProtocolVersion: 5, + PacketID: pk.PacketID, + } + var ab bytes.Buffer + _ = ack.PubackEncode(&ab) + writeMu.Lock() + _, _ = w.Write(ab.Bytes()) + writeMu.Unlock() + } + }() + + go func() { + tick := time.NewTicker(2 * time.Millisecond) + defer tick.Stop() + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Pingreq}, + ProtocolVersion: 5, + } + var buf bytes.Buffer + _ = pk.PingreqEncode(&buf) + ping := append([]byte(nil), buf.Bytes()...) + for { + select { + case <-stop: + return + case <-tick.C: + writeMu.Lock() + _, _ = w.Write(ping) + writeMu.Unlock() + } + } + }() + + payload := []byte(`{"v":1,"type":"resp","rid":"x"}`) + pubCtx, cancel := context.WithCancel(context.Background()) + defer cancel() + var wg sync.WaitGroup + for i := 0; i < 8; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for pubCtx.Err() == nil { + _ = b.PublishDown(pubCtx, endpoint, "", payload, port.PublishOpts{QoS: 1}) + } + }() + } + + runFor := 5 * time.Second + watch := 2 * time.Second + start := time.Now() + last := received.Load() + lastChange := time.Now() + for time.Since(start) < runFor { + time.Sleep(50 * time.Millisecond) + n := received.Load() + if n > last { + last = n + lastChange = time.Now() + } + if time.Since(lastChange) > watch { + buf := make([]byte, 1<<20) + nstack := runtime.Stack(buf, true) + halt() + cancel() + _ = w.Close() + t.Fatalf("progress stalled at %d after %s\n%s", last, time.Since(lastChange), buf[:nstack]) + } + } + cancel() + wg.Wait() + halt() + if last < 100 { + t.Fatalf("too few publishes delivered: %d", last) + } + + _ = w.Close() + select { + case <-clientDone: + case <-time.After(3 * time.Second): + } +} + +func readExactPacket(t *testing.T, conn net.Conn, wantType byte, timeout time.Duration) { + t.Helper() + _ = conn.SetReadDeadline(time.Now().Add(timeout)) + hdr := make([]byte, 1) + if _, err := io.ReadFull(conn, hdr); err != nil { + t.Fatal(err) + } + if hdr[0]>>4 != wantType { + t.Fatalf("want packet type %d got %d", wantType, hdr[0]>>4) + } + rem, err := readRemainingLengthConn(conn) + if err != nil { + t.Fatal(err) + } + body := make([]byte, rem) + if _, err := io.ReadFull(conn, body); err != nil { + t.Fatal(err) + } +} + +func readRemainingLengthConn(r io.Reader) (int, error) { + var mul uint32 = 1 + var value uint32 + for i := 0; i < 4; i++ { + var b [1]byte + if _, err := io.ReadFull(r, b[:]); err != nil { + return 0, err + } + value += uint32(b[0]&127) * mul + if b[0]&128 == 0 { + return int(value), nil + } + mul *= 128 + } + return 0, io.ErrUnexpectedEOF +} diff --git a/internal/broker/broker_test.go b/internal/broker/broker_test.go index 69a66c2..bdabdb3 100644 --- a/internal/broker/broker_test.go +++ b/internal/broker/broker_test.go @@ -212,6 +212,11 @@ func connectAndSubscribe(t *testing.T, w net.Conn, endpoint string, maxPacket ui } func writeConnect(t *testing.T, w net.Conn, endpoint string, keepalive uint16, maxPacket uint32) { + t.Helper() + writeConnectFull(t, w, endpoint, keepalive, maxPacket, 0) +} + +func writeConnectFull(t *testing.T, w net.Conn, endpoint string, keepalive uint16, maxPacket uint32, receiveMax uint16) { t.Helper() pk := packets.Packet{ FixedHeader: packets.FixedHeader{Type: packets.Connect}, @@ -228,6 +233,7 @@ func writeConnect(t *testing.T, w net.Conn, endpoint string, keepalive uint16, m }, Properties: packets.Properties{ MaximumPacketSize: maxPacket, + ReceiveMaximum: receiveMax, }, } var buf bytes.Buffer diff --git a/internal/broker/hooks.go b/internal/broker/hooks.go index 758da07..d63f239 100644 --- a/internal/broker/hooks.go +++ b/internal/broker/hooks.go @@ -67,6 +67,16 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { cl.State.ServerKeepalive = true } + // B-01:绕开 mochi 发送配额路径(NextImmediate 递归读锁 + PUBACK 配额泄漏)。 + // ParseConnect 已按客户端 Receive Maximum 设过 sendQuota;此处一律置 0。 + if cl.State.Inflight != nil { + cl.State.Inflight.ResetSendQuota(0) + } + if rm := pk.Properties.ReceiveMaximum; rm > 0 && rm < 256 { + h.b.log.Warn("client receive maximum below 256; server ignores MQTT send quota", + "endpoint", endpointID, "receive_maximum", rm) + } + res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP) if err != nil { st.authErr = err From b5b63ed070383bd6dd3f3a7b88f5e4be83ceacb3 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:05:42 +0800 Subject: [PATCH 2/4] =?UTF-8?q?fix:=20=E5=A4=A7=E5=B8=A7=E5=90=8D=E9=A2=9D?= =?UTF-8?q?=E6=8C=89=20PacketID=20=E5=9C=A8=20PUBACK=20=E4=B8=8E=E6=96=AD?= =?UTF-8?q?=E7=BA=BF=E6=97=B6=E5=BD=92=E8=BF=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 9 ++ internal/broker/b02_large_test.go | 213 ++++++++++++++++++++++++++++++ internal/broker/broker.go | 100 ++++++++++---- internal/broker/hooks.go | 40 +++++- 4 files changed, 331 insertions(+), 31 deletions(-) create mode 100644 internal/broker/b02_large_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 2f4f30c..45da5b7 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1180,3 +1180,12 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 - 原因:mochi v2.7.9 在 `sendQuota>0` 时走 `NextImmediate` 递归读锁,并可因补发后删除 inflight 泄漏配额;已验证置 0 绕开整条路径。 - 备选方案:fork 修补 mochi(只修递归读锁仍观察到停滞)。 - 影响:服务端不再执行客户端 Receive Maximum;裸设备若带过小的 Receive Maximum,实际在途可能超过该值。 + +### 复审修复 B-02 + +- 日期:2026-09-30 +- 原条款:DEVELOPMENT 7.5 / DEVIATIONS N1/N2 第 4 条:大帧名额在 PUBACK、丢弃、断线时归还。Gitea #9。 +- 实际做法:`OnQosPublish` 按 PacketID 记下超过 64KiB 的出站包;`OnQosComplete`/`OnQosDropped`/断线按 ID 归还。获取名额最多等 5 秒,超时返回 `ErrLargeFrameTimeout`。Publish 未产生 inflight(无订阅者、队列丢弃)时立即归还。不采用「发布完成即归还」。 +- 原因:mochi 传给 `OnQosComplete` 的是 PUBACK,没有载荷,旧实现从未归还。 +- 备选方案:发布后立即归还(会把卡死点挪到消息包那份名额)。 +- 影响:只在 broker 保留一份全局 64 名额;确认超时仍由消息线踢线/清标记触发断线归还。 diff --git a/internal/broker/b02_large_test.go b/internal/broker/b02_large_test.go new file mode 100644 index 0000000..f353b73 --- /dev/null +++ b/internal/broker/b02_large_test.go @@ -0,0 +1,213 @@ +package broker + +import ( + "bytes" + "context" + "io" + "net" + "sync" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" + "github.com/mochi-mqtt/server/v2/packets" +) + +func TestLargeFrameQuotaReleasedOnPuback(t *testing.T) { + b, w, done := startTCPClient(t, "ep-large-ack") + defer func() { _ = b.Close() }() + defer func() { + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + }() + + writeConnect(t, w, "ep-large-ack", 30, 0) + readExactPacket(t, w, packets.Connack, 3*time.Second) + writeSubscribe(t, w, downTopic("ep-large-ack")) + readExactPacket(t, w, packets.Suback, 3*time.Second) + waitSession(t, b, "ep-large-ack") + + stop := make(chan struct{}) + defer close(stop) + var writeMu sync.Mutex + go autoPuback(w, stop, &writeMu) + + payload := bytes.Repeat([]byte("x"), 70*1024) + for i := 0; i < 65; i++ { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + err := b.PublishDown(ctx, "ep-large-ack", "", payload, port.PublishOpts{QoS: 1}) + cancel() + if err != nil { + t.Fatalf("publish %d: %v", i+1, err) + } + } + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + if len(b.largeSem) == 0 { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("slots still held: %d", len(b.largeSem)) +} + +func TestLargeFrameQuotaReleasedOnDisconnect(t *testing.T) { + b, w, done := startTCPClient(t, "ep-large-disc") + defer func() { _ = b.Close() }() + + writeConnect(t, w, "ep-large-disc", 30, 0) + readExactPacket(t, w, packets.Connack, 3*time.Second) + writeSubscribe(t, w, downTopic("ep-large-disc")) + readExactPacket(t, w, packets.Suback, 3*time.Second) + waitSession(t, b, "ep-large-disc") + + go func() { + buf := make([]byte, 32*1024) + for { + _, err := w.Read(buf) + if err != nil { + return + } + } + }() + + payload := bytes.Repeat([]byte("y"), 70*1024) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + if err := b.PublishDown(ctx, "ep-large-disc", "", payload, port.PublishOpts{QoS: 1}); err != nil { + cancel() + t.Fatal(err) + } + cancel() + if len(b.largeSem) == 0 { + t.Fatal("expected a held large slot before disconnect") + } + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + t.Fatal("client attach did not return") + } + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + if len(b.largeSem) == 0 { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("slots after disconnect: %d", len(b.largeSem)) +} + +func TestLargeFrameQuotaReleasedWithoutSubscriber(t *testing.T) { + b, w, done := startTCPClient(t, "ep-large-nosub") + defer func() { _ = b.Close() }() + defer func() { + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + }() + + writeConnect(t, w, "ep-large-nosub", 30, 0) + readExactPacket(t, w, packets.Connack, 3*time.Second) + waitSession(t, b, "ep-large-nosub") + + payload := bytes.Repeat([]byte("z"), 70*1024) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + err := b.PublishDown(ctx, "ep-large-nosub", "", payload, port.PublishOpts{QoS: 1}) + cancel() + if err != nil { + t.Fatal(err) + } + if n := len(b.largeSem); n != 0 { + t.Fatalf("held slots without subscriber: %d", n) + } +} + +func startTCPClient(t *testing.T, _ string) (*Broker, net.Conn, chan struct{}) { + t.Helper() + b, err := New(Options{Authenticator: AllowAuthenticator{}}) + if err != nil { + t.Fatal(err) + } + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = ln.Close() }) + done := make(chan struct{}) + go func() { + defer close(done) + c, accErr := ln.Accept() + if accErr != nil { + return + } + _ = b.AttachTCP(c) + }() + w, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + t.Fatal(err) + } + return b, w, done +} + +func waitSession(t *testing.T, b *Broker, endpoint string) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + if _, ok := b.ConnInfoOf(endpoint); ok { + return + } + time.Sleep(5 * time.Millisecond) + } + t.Fatal("session not established") +} + +func autoPuback(w net.Conn, stop <-chan struct{}, writeMu *sync.Mutex) { + for { + select { + case <-stop: + return + default: + } + _ = 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 + } + if hdr[0]>>4 != packets.Publish { + continue + } + qos := (hdr[0] >> 1) & 0x3 + if qos == 0 { + continue + } + 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 { + continue + } + ack := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Puback}, + ProtocolVersion: 5, + PacketID: pk.PacketID, + } + var ab bytes.Buffer + _ = ack.PubackEncode(&ab) + writeMu.Lock() + _, _ = w.Write(ab.Bytes()) + writeMu.Unlock() + } +} diff --git a/internal/broker/broker.go b/internal/broker/broker.go index 0ecfdd5..81e011b 100644 --- a/internal/broker/broker.go +++ b/internal/broker/broker.go @@ -34,6 +34,11 @@ var ErrPayloadTooLarge = errors.New("broker: payload exceeds client limit") // ErrNoConnection 目标端没有当前连接。 var ErrNoConnection = errors.New("broker: no active connection") +// ErrLargeFrameTimeout 全局大帧名额在有界等待内拿不到。 +var ErrLargeFrameTimeout = errors.New("broker: large frame quota timeout") + +const largeAcquireWait = 5 * time.Second + // AuthResult 是登录校验结论(N3 实现真实逻辑;N2 默认拒绝)。 type AuthResult struct { OK bool @@ -108,7 +113,8 @@ type connState struct { sessionToken string handshook bool subscribedDown bool - largeHeld int + largePIDs map[uint16]struct{} + largePending int metricsCounted bool mu sync.Mutex @@ -219,54 +225,98 @@ func (b *Broker) PublishDown(ctx context.Context, endpointID string, connID port } topic := downTopic(endpointID) large := len(payload) > largeFrameBytes - if large { - select { - case b.largeSem <- struct{}{}: - case <-ctx.Done(): - return ctx.Err() + if err := b.acquireLarge(ctx); err != nil { + return err } st.mu.Lock() - st.largeHeld++ + st.largePending++ st.mu.Unlock() } if err := b.server.Publish(topic, payload, false, qos); err != nil { if large { - b.releaseOneLarge(st) + b.finishLargePublish(st) } return err } - if large && qos == 0 { - b.releaseOneLarge(st) + if large { + b.finishLargePublish(st) } return nil } -func (b *Broker) releaseOneLarge(st *connState) { +func (b *Broker) acquireLarge(ctx context.Context) error { + timer := time.NewTimer(largeAcquireWait) + defer timer.Stop() + select { + case b.largeSem <- struct{}{}: + return nil + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return ErrLargeFrameTimeout + } +} + +func (b *Broker) releaseLargeSlot() { + select { + case <-b.largeSem: + default: + } +} + +func (b *Broker) finishLargePublish(st *connState) { + b.reconcileLargeInflight(st) st.mu.Lock() - if st.largeHeld > 0 { - st.largeHeld-- - st.mu.Unlock() - select { - case <-b.largeSem: - default: - } - return + n := st.largePending + st.largePending = 0 + st.mu.Unlock() + for i := 0; i < n; i++ { + b.releaseLargeSlot() + } +} + +func (b *Broker) releaseLargePID(st *connState, id uint16) { + st.mu.Lock() + _, ok := st.largePIDs[id] + if ok { + delete(st.largePIDs, id) } st.mu.Unlock() + if ok { + b.releaseLargeSlot() + } +} + +func (b *Broker) reconcileLargeInflight(st *connState) { + if st == nil { + return + } + st.mu.Lock() + ids := make([]uint16, 0, len(st.largePIDs)) + for id := range st.largePIDs { + ids = append(ids, id) + } + st.mu.Unlock() + for _, id := range ids { + if st.client != nil { + if _, ok := st.client.State.Inflight.Get(id); ok { + continue + } + } + b.releaseLargePID(st, id) + } } func (b *Broker) releaseAllLarge(st *connState) { st.mu.Lock() - n := st.largeHeld - st.largeHeld = 0 + n := len(st.largePIDs) + st.largePending + st.largePIDs = nil + st.largePending = 0 st.mu.Unlock() for i := 0; i < n; i++ { - select { - case <-b.largeSem: - default: - } + b.releaseLargeSlot() } } diff --git a/internal/broker/hooks.go b/internal/broker/hooks.go index d63f239..3638703 100644 --- a/internal/broker/hooks.go +++ b/internal/broker/hooks.go @@ -25,7 +25,9 @@ func (h *nixHook) Provides(b byte) bool { mqtt.OnPublishDropped, mqtt.OnSessionEstablished, mqtt.OnDisconnect, + mqtt.OnQosPublish, mqtt.OnQosComplete, + mqtt.OnQosDropped, mqtt.OnSubscribed, }, []byte{b}) } @@ -169,13 +171,13 @@ func (h *nixHook) OnSubscribed(cl *mqtt.Client, pk packets.Packet, reasonCodes [ func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) { h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload)) - if h.b.onDrop == nil { - return - } h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() - if st == nil { + if st != nil { + h.b.reconcileLargeInflight(st) + } + if h.b.onDrop == nil || st == nil { return } h.b.onDrop(context.Background(), st.endpointID, st.connID, append([]byte(nil), pk.Payload...)) @@ -263,7 +265,7 @@ func (h *nixHook) noteConnectionClose(st *connState) { st.metricsCounted = false } -func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) { +func (h *nixHook) OnQosPublish(cl *mqtt.Client, pk packets.Packet, _ int64, _ int) { if len(pk.Payload) <= largeFrameBytes { return } @@ -273,5 +275,31 @@ func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) { if st == nil { return } - h.b.releaseOneLarge(st) + st.mu.Lock() + if st.largePending > 0 { + st.largePending-- + } + if st.largePIDs == nil { + st.largePIDs = make(map[uint16]struct{}) + } + st.largePIDs[pk.PacketID] = struct{}{} + st.mu.Unlock() +} + +func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) { + h.releaseLargeByPacketID(cl, pk.PacketID) +} + +func (h *nixHook) OnQosDropped(cl *mqtt.Client, pk packets.Packet) { + h.releaseLargeByPacketID(cl, pk.PacketID) +} + +func (h *nixHook) releaseLargeByPacketID(cl *mqtt.Client, id uint16) { + h.b.connsMu.RLock() + st := h.b.byClient[cl] + h.b.connsMu.RUnlock() + if st == nil { + return + } + h.b.releaseLargePID(st, id) } From 42160720bd49d1f456e5e3a3b472ce92318a9dad Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:08:03 +0800 Subject: [PATCH 3/4] =?UTF-8?q?fix:=20=E8=AE=A4=E8=AF=81=E5=A4=B1=E8=B4=A5?= =?UTF-8?q?=E4=B8=8D=E6=B3=84=E6=BC=8F=E8=BF=9E=E6=8E=A5=E8=A1=A8=E5=B9=B6?= =?UTF-8?q?=E8=84=B1=E6=95=8F=20mochi=20=E6=95=B4=E5=8C=85=E6=97=A5?= =?UTF-8?q?=E5=BF=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 18 +++ internal/broker/b05_b07_test.go | 270 ++++++++++++++++++++++++++++++++ internal/broker/broker.go | 62 +++++++- internal/broker/hooks.go | 18 ++- internal/broker/log.go | 88 +++++++++++ 5 files changed, 442 insertions(+), 14 deletions(-) create mode 100644 internal/broker/b05_b07_test.go create mode 100644 internal/broker/log.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 45da5b7..bbea34b 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1189,3 +1189,21 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 - 原因:mochi 传给 `OnQosComplete` 的是 PUBACK,没有载荷,旧实现从未归还。 - 备选方案:发布后立即归还(会把卡死点挪到消息包那份名额)。 - 影响:只在 broker 保留一份全局 64 名额;确认超时仍由消息线踢线/清标记触发断线归还。 + +### 复审修复 B-05 + +- 日期:2026-09-30 +- 原条款:DEVELOPMENT 第 5 节连接表;Gitea #12。 +- 实际做法:`OnConnect` 只在认证通过时写入 `byClient`/`byConnID`;拒绝与内部错误不登记。`connState` 增加 `established` 与 `createdAt`,每分钟清扫未建立且已关闭超过 1 分钟的条目。按连接代号查找改为 O(1)。 +- 原因:mochi 在认证失败路径不调用 `OnDisconnect`,旧实现会永久泄漏。 +- 备选方案:失败路径也登记再在 Authenticate 返回 false 时删除(仍覆盖不了 CONNACK 失败)。 +- 影响:失败连接不再占用查找路径;行为对客户端不变(仍回 0x86 或不回 CONNACK)。 + +### 复审修复 B-07 + +- 日期:2026-09-30 +- 原条款:PRD §8 日志无正文、无密码、无令牌。Gitea #14。 +- 实际做法:`broker.New` 给 mochi 包一层 slog.Handler,把 `packets.Packet` / `*packets.Packet` 换成类型、QoS、包号、主题、正文长度。 +- 原因:默认 info 下第二个 CONNECT、3.1.1 发到错误主题等会把整包写入 JSON 日志。 +- 备选方案:改 mochi 日志调用点(需 fork)。 +- 影响:排障时看不到载荷与密码,只见摘要。 diff --git a/internal/broker/b05_b07_test.go b/internal/broker/b05_b07_test.go new file mode 100644 index 0000000..250be2a --- /dev/null +++ b/internal/broker/b05_b07_test.go @@ -0,0 +1,270 @@ +package broker + +import ( + "bytes" + "context" + "encoding/base64" + "io" + "log/slog" + "net" + "strings" + "testing" + "time" + + "github.com/mochi-mqtt/server/v2/packets" +) + +func TestFailedAuthDoesNotLeakConnTable(t *testing.T) { + secret := "s3cret-token-xyz" + var logBuf bytes.Buffer + log := slog.New(slog.NewTextHandler(&logBuf, &slog.HandlerOptions{Level: slog.LevelDebug})) + + b, err := New(Options{Authenticator: RejectAuthenticator{}, Logger: log}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b.Close() }() + + const n = 200 + for i := 0; i < n; i++ { + dialFailedCONNECT(t, b, func(w net.Conn) { + writeConnect(t, w, "ep-rej", 30, 0) + }) + } + for i := 0; i < n; i++ { + dialFailedCONNECT(t, b, func(w net.Conn) { + writeConnectMismatch(t, w) + }) + } + b2, err := New(Options{Authenticator: &errAuthenticator{err: context.DeadlineExceeded}}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b2.Close() }() + for i := 0; i < n; i++ { + dialFailedCONNECT(t, b2, func(w net.Conn) { + writeConnect(t, w, "ep-err", 30, 0) + }) + } + if got := len(b.byClient); got != 0 { + t.Fatalf("reject/mismatch leaked %d", got) + } + if got := len(b2.byClient); got != 0 { + t.Fatalf("internal error leaked %d", got) + } + + // B-07:拒绝路径的 mochi 日志不能带密码 + b3, err := New(Options{Authenticator: AllowAuthenticator{}, Logger: log}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b3.Close() }() + r, w := net.Pipe() + done := make(chan struct{}) + go func() { + defer close(done) + _ = b3.AttachTCP(r) + }() + writeConnectWithPassword(t, w, "ep-log", secret) + readExactPacket(t, w, packets.Connack, 3*time.Second) + writeConnectWithPassword(t, w, "ep-log", secret) // 同一连接第二个 CONNECT + _ = w.Close() + select { + case <-done: + case <-time.After(2 * time.Second): + } + out := logBuf.String() + if strings.Contains(out, secret) { + t.Fatalf("log contains password: %s", out) + } + if strings.Contains(out, base64.StdEncoding.EncodeToString([]byte(secret))) { + t.Fatalf("log contains password base64: %s", out) + } +} + +func TestSweepUnestablishedClosedConn(t *testing.T) { + b, err := New(Options{Authenticator: AllowAuthenticator{}}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b.Close() }() + + r, w := net.Pipe() + done := make(chan struct{}) + go func() { + defer close(done) + _ = b.AttachTCP(r) + }() + writeConnect(t, w, "ep-sweep", 30, 0) + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + b.sweepUnestablished(0) + if got := len(b.byClient); got != 0 { + t.Fatalf("after sweep byClient=%d", got) + } +} + +func TestLookupByConnIDIndependentOfFailedConns(t *testing.T) { + b, err := New(Options{Authenticator: AllowAuthenticator{}}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b.Close() }() + r, w := net.Pipe() + done := make(chan struct{}) + go func() { + defer close(done) + _ = b.AttachTCP(r) + }() + connectAndSubscribe(t, w, "ep-ok", 0) + waitSession(t, b, "ep-ok") + info, ok := b.ConnInfoOf("ep-ok") + if !ok { + t.Fatal("missing session") + } + st := b.lookupConn("ep-ok", info.ConnID) + if st == nil { + t.Fatal("lookup by conn id") + } + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } +} + +func dialFailedCONNECT(t *testing.T, b *Broker, write func(net.Conn)) { + t.Helper() + r, w := net.Pipe() + done := make(chan struct{}) + go func() { + defer close(done) + _ = b.AttachTCP(r) + }() + write(w) + _ = w.Close() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("attach did not return") + } +} + +func writeConnectMismatch(t *testing.T, w net.Conn) { + t.Helper() + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Connect}, + ProtocolVersion: 5, + Connect: packets.ConnectParams{ + ProtocolName: []byte("MQTT"), + Clean: true, + ClientIdentifier: "id-a", + Keepalive: 30, + UsernameFlag: true, + Username: []byte("id-b"), + PasswordFlag: true, + Password: []byte("nope"), + }, + } + var buf bytes.Buffer + if err := pk.ConnectEncode(&buf); err != nil { + t.Fatal(err) + } + if _, err := w.Write(buf.Bytes()); err != nil { + t.Fatal(err) + } +} + +func writeConnectWithPassword(t *testing.T, w net.Conn, endpoint, password string) { + t.Helper() + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Connect}, + ProtocolVersion: 5, + Connect: packets.ConnectParams{ + ProtocolName: []byte("MQTT"), + Clean: true, + ClientIdentifier: endpoint, + Keepalive: 30, + UsernameFlag: true, + Username: []byte(endpoint), + PasswordFlag: true, + Password: []byte(password), + }, + } + var buf bytes.Buffer + if err := pk.ConnectEncode(&buf); err != nil { + t.Fatal(err) + } + if _, err := w.Write(buf.Bytes()); err != nil { + t.Fatal(err) + } +} + +func TestMQTT311UnauthorizedPublishOmitsPayloadInLogs(t *testing.T) { + var logBuf bytes.Buffer + log := slog.New(slog.NewTextHandler(&logBuf, &slog.HandlerOptions{Level: slog.LevelDebug})) + b, err := New(Options{Authenticator: AllowAuthenticator{}, Logger: log}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b.Close() }() + + r, w := net.Pipe() + done := make(chan struct{}) + go func() { + defer close(done) + _ = b.AttachTCP(r) + }() + pk := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Connect}, + ProtocolVersion: 4, + Connect: packets.ConnectParams{ + ProtocolName: []byte("MQTT"), + Clean: true, + ClientIdentifier: "ep311", + Keepalive: 30, + UsernameFlag: true, + Username: []byte("ep311"), + PasswordFlag: true, + Password: []byte("test"), + }, + } + var buf bytes.Buffer + if err := pk.ConnectEncode(&buf); err != nil { + t.Fatal(err) + } + if _, err := w.Write(buf.Bytes()); err != nil { + t.Fatal(err) + } + _ = w.SetReadDeadline(time.Now().Add(3 * time.Second)) + raw := make([]byte, 256) + if _, err := io.ReadAtLeast(w, raw, 2); err != nil { + t.Fatal(err) + } + body := []byte(`{"talk_password":"super-secret-body"}`) + pub := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Publish, Qos: 1}, + ProtocolVersion: 4, + TopicName: "nix/c/other/up", + PacketID: 7, + Payload: body, + } + buf.Reset() + if err := pub.PublishEncode(&buf); err != nil { + t.Fatal(err) + } + _, _ = w.Write(buf.Bytes()) + time.Sleep(50 * time.Millisecond) + _ = w.Close() + select { + case <-done: + case <-time.After(2 * time.Second): + } + out := logBuf.String() + if strings.Contains(out, "super-secret-body") { + t.Fatalf("log contains publish payload: %s", out) + } +} diff --git a/internal/broker/broker.go b/internal/broker/broker.go index 81e011b..91accc6 100644 --- a/internal/broker/broker.go +++ b/internal/broker/broker.go @@ -92,6 +92,8 @@ type Broker struct { connsMu sync.RWMutex current map[string]*connState byClient map[*mqtt.Client]*connState + byConnID map[port.ConnID]*connState + closedCh chan struct{} queuesMu sync.Mutex queues map[string]*uplinkQueue @@ -116,6 +118,8 @@ type connState struct { largePIDs map[uint16]struct{} largePending int metricsCounted bool + established bool + createdAt time.Time mu sync.Mutex handshakeTimer *time.Timer @@ -135,6 +139,7 @@ func New(opts Options) (*Broker, error) { if log == nil { log = slog.Default() } + log = slog.New(newRedactHandler(log.Handler())) caps := mqtt.NewDefaultServerCapabilities() caps.MaximumClients = maxClients @@ -165,6 +170,8 @@ func New(opts Options) (*Broker, error) { metrics: opts.Metrics, current: make(map[string]*connState), byClient: make(map[*mqtt.Client]*connState), + byConnID: make(map[port.ConnID]*connState), + closedCh: make(chan struct{}), queues: make(map[string]*uplinkQueue), largeSem: make(chan struct{}, largeFrameSlots), } @@ -175,6 +182,7 @@ func New(opts Options) (*Broker, error) { if err := srv.Serve(); err != nil { return nil, err } + go b.sweepLoop() return b, nil } @@ -186,6 +194,11 @@ func (b *Broker) Close() error { if b.closed.Swap(true) { return nil } + select { + case <-b.closedCh: + default: + close(b.closedCh) + } b.queuesMu.Lock() for _, q := range b.queues { q.close() @@ -345,10 +358,9 @@ func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState { b.connsMu.RLock() defer b.connsMu.RUnlock() if connID != "" { - for _, st := range b.byClient { - if st.endpointID == endpointID && st.connID == connID { - return st - } + st := b.byConnID[connID] + if st != nil && st.endpointID == endpointID { + return st } return nil } @@ -462,12 +474,46 @@ func (b *Broker) CurrentConnID(endpointID string) (port.ConnID, bool) { func (b *Broker) connStateOf(endpointID string, connID port.ConnID) *connState { b.connsMu.RLock() defer b.connsMu.RUnlock() - for _, st := range b.byClient { - if st.endpointID == endpointID && st.connID == connID { - return st + st := b.byConnID[connID] + if st == nil || st.endpointID != endpointID { + return nil + } + return st +} + +func (b *Broker) sweepLoop() { + tick := time.NewTicker(time.Minute) + defer tick.Stop() + for { + select { + case <-tick.C: + b.sweepUnestablished(time.Minute) + case <-b.closedCh: + return + } + } +} + +func (b *Broker) sweepUnestablished(minAge time.Duration) { + now := time.Now() + b.connsMu.Lock() + defer b.connsMu.Unlock() + for cl, st := range b.byClient { + if st.established { + continue + } + if cl != nil && !cl.Closed() { + continue + } + if minAge > 0 && now.Sub(st.createdAt) < minAge { + continue + } + delete(b.byClient, cl) + delete(b.byConnID, st.connID) + if b.current[st.endpointID] == st { + delete(b.current, st.endpointID) } } - return nil } func (b *Broker) hasDownSub(st *connState) bool { diff --git a/internal/broker/hooks.go b/internal/broker/hooks.go index 3638703..b6795fb 100644 --- a/internal/broker/hooks.go +++ b/internal/broker/hooks.go @@ -3,6 +3,7 @@ package broker import ( "bytes" "context" + "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" mqtt "github.com/mochi-mqtt/server/v2" @@ -47,12 +48,11 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { remoteIP: remoteIP, client: cl, maxPacketSize: pk.Properties.MaximumPacketSize, + createdAt: time.Now(), } // ClientID、Username 都必须等于端编号 if clientID == "" || endpointID == "" || clientID != endpointID { - st.authOK = false - h.rememberPending(cl, st) return nil } @@ -81,11 +81,12 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP) if err != nil { - st.authErr = err - h.rememberPending(cl, st) - return err // mochi 不回 CONNACK,直接断开 + return err // mochi 不回 CONNACK,直接断开;不登记连接表 } - st.authOK = res.OK + if !res.OK { + return nil + } + st.authOK = true st.sessionToken = res.SessionToken h.rememberPending(cl, st) return nil @@ -94,6 +95,7 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { func (h *nixHook) rememberPending(cl *mqtt.Client, st *connState) { h.b.connsMu.Lock() h.b.byClient[cl] = st + h.b.byConnID[st.connID] = st h.b.connsMu.Unlock() } @@ -188,6 +190,7 @@ func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) { st := h.b.byClient[cl] if st != nil { h.b.current[st.endpointID] = st + st.established = true } h.b.connsMu.Unlock() if st == nil { @@ -209,6 +212,9 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) { h.b.connsMu.Lock() st := h.b.byClient[cl] delete(h.b.byClient, cl) + if st != nil { + delete(h.b.byConnID, st.connID) + } isCurrent := false if st != nil && h.b.current[st.endpointID] == st { delete(h.b.current, st.endpointID) diff --git a/internal/broker/log.go b/internal/broker/log.go new file mode 100644 index 0000000..b32efda --- /dev/null +++ b/internal/broker/log.go @@ -0,0 +1,88 @@ +package broker + +import ( + "context" + "log/slog" + + "github.com/mochi-mqtt/server/v2/packets" +) + +type redactHandler struct { + inner slog.Handler +} + +func newRedactHandler(inner slog.Handler) slog.Handler { + if inner == nil { + inner = slog.Default().Handler() + } + return &redactHandler{inner: inner} +} + +func (h *redactHandler) Enabled(ctx context.Context, level slog.Level) bool { + return h.inner.Enabled(ctx, level) +} + +func (h *redactHandler) Handle(ctx context.Context, r slog.Record) error { + rec := slog.NewRecord(r.Time, r.Level, r.Message, r.PC) + r.Attrs(func(a slog.Attr) bool { + rec.AddAttrs(redactSlogAttr(a)) + return true + }) + return h.inner.Handle(ctx, rec) +} + +func (h *redactHandler) WithAttrs(attrs []slog.Attr) slog.Handler { + out := make([]slog.Attr, len(attrs)) + for i, a := range attrs { + out[i] = redactSlogAttr(a) + } + return &redactHandler{inner: h.inner.WithAttrs(out)} +} + +func (h *redactHandler) WithGroup(name string) slog.Handler { + return &redactHandler{inner: h.inner.WithGroup(name)} +} + +func redactSlogAttr(a slog.Attr) slog.Attr { + a.Value = a.Value.Resolve() + switch v := a.Value.Any().(type) { + case packets.Packet: + return slog.Any(a.Key, summarizePacket(v)) + case *packets.Packet: + if v == nil { + return a + } + return slog.Any(a.Key, summarizePacket(*v)) + } + if a.Value.Kind() == slog.KindGroup { + group := a.Value.Group() + out := make([]slog.Attr, len(group)) + for i, g := range group { + out[i] = redactSlogAttr(g) + } + return slog.Attr{Key: a.Key, Value: slog.GroupValue(out...)} + } + return a +} + +type mqttPacketLog struct { + Type string `json:"type"` + QoS byte `json:"qos"` + PacketID uint16 `json:"packet_id"` + Topic string `json:"topic,omitempty"` + PayloadLen int `json:"payload_len"` +} + +func summarizePacket(pk packets.Packet) mqttPacketLog { + name := packets.PacketNames[pk.FixedHeader.Type] + if name == "" { + name = "unknown" + } + return mqttPacketLog{ + Type: name, + QoS: pk.FixedHeader.Qos, + PacketID: pk.PacketID, + Topic: pk.TopicName, + PayloadLen: len(pk.Payload), + } +} From 8b845843d6824d96e21835eeb637c97c26cbb3f1 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:24:15 +0800 Subject: [PATCH 4/4] =?UTF-8?q?fix:=20=E5=AE=8C=E6=88=90=20broker=20?= =?UTF-8?q?=E5=A4=8D=E5=AE=A1=20B-03=20=E8=87=B3=20B-12?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 每连接异步下发与背压、写出后断开、校验当前连接与订阅、生命周期串行、登录条件更新、闲置按在线计、认证超时并发与 Shutdown 0x8B。 --- cmd/nixmsg/serve.go | 6 +- cmd/nixmsg/uplink.go | 58 ++++++-- docs/DEVIATIONS.md | 72 +++++++++ internal/broker/authn.go | 128 ++++++++++++++-- internal/broker/authn_idle_test.go | 123 ++++++++++++++++ internal/broker/b02_large_test.go | 22 ++- internal/broker/b03_downlink_test.go | 212 +++++++++++++++++++++++++++ internal/broker/b04_b06_b08_test.go | 128 ++++++++++++++++ internal/broker/broker.go | 177 ++++++++++++++++------ internal/broker/downlink.go | 192 ++++++++++++++++++++++++ internal/broker/hooks.go | 32 +++- internal/broker/session.go | 38 +++-- 12 files changed, 1099 insertions(+), 89 deletions(-) create mode 100644 internal/broker/authn_idle_test.go create mode 100644 internal/broker/b03_downlink_test.go create mode 100644 internal/broker/b04_b06_b08_test.go create mode 100644 internal/broker/downlink.go diff --git a/cmd/nixmsg/serve.go b/cmd/nixmsg/serve.go index 0fa3336..51cd1ec 100644 --- a/cmd/nixmsg/serve.go +++ b/cmd/nixmsg/serve.go @@ -153,7 +153,7 @@ func runServe(ctx context.Context, cfg config.Config) error { Sessions: sessionTokens, MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds), Logger: slog.Default(), - ConnControl: brk, + ConnControl: nil, // B-04:踢线走 Session 钩子,避免 identity 20ms 异步 Disconnect Downlink: brk, ClientIP: func(r *http.Request) string { return httpx.ClientIP(r, trustedNets) @@ -310,6 +310,10 @@ func runServe(ctx context.Context, cfg config.Config) error { <-ctx.Done() loopCancel() + // B-08:先对 MQTT 连接发 0x8B。HTTP Shutdown 与监听器完整停机顺序见 L-03。 + shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second) + _ = brk.Shutdown(shutCtx) + shutCancel() _ = lnSrv.Close() drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second) defer drainCancel() diff --git a/cmd/nixmsg/uplink.go b/cmd/nixmsg/uplink.go index 1f429ff..6f4ab57 100644 --- a/cmd/nixmsg/uplink.go +++ b/cmd/nixmsg/uplink.go @@ -5,12 +5,14 @@ import ( "encoding/json" "errors" "log/slog" + "sync" "git.asio.asia/nixevol/NixMsg/internal/app/group" "git.asio.asia/nixevol/NixMsg/internal/app/identity" "git.asio.asia/nixevol/NixMsg/internal/app/message" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/presence" + "git.asio.asia/nixevol/NixMsg/internal/broker" "git.asio.asia/nixevol/NixMsg/internal/metrics" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) @@ -25,9 +27,31 @@ type appUplink struct { down port.Downlink log *slog.Logger metrics *metrics.Registry + + lifeMu sync.Mutex + lifeLocks map[string]*sync.Mutex + hsMu sync.Mutex + handshake map[port.ConnID]string // 已 hello 的连接代号 → 端编号 +} + +func (u *appUplink) epLife(endpointID string) *sync.Mutex { + u.lifeMu.Lock() + defer u.lifeMu.Unlock() + if u.lifeLocks == nil { + u.lifeLocks = make(map[string]*sync.Mutex) + } + m := u.lifeLocks[endpointID] + if m == nil { + m = &sync.Mutex{} + u.lifeLocks[endpointID] = m + } + return m } func (u *appUplink) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error { + lk := u.epLife(conn.EndpointID) + lk.Lock() + defer lk.Unlock() u.conns.Set(conn.EndpointID, message.LiveConn{ ConnID: conn.ConnID, MaxPacketSize: conn.MaxPacketSize, @@ -36,6 +60,15 @@ func (u *appUplink) OnSessionEstablished(ctx context.Context, conn port.ConnInfo } func (u *appUplink) OnHandshakeComplete(ctx context.Context, hs port.HandshakeInfo) error { + lk := u.epLife(hs.EndpointID) + lk.Lock() + defer lk.Unlock() + u.hsMu.Lock() + if u.handshake == nil { + u.handshake = make(map[port.ConnID]string) + } + u.handshake[hs.ConnID] = hs.EndpointID + u.hsMu.Unlock() live := message.LiveConn{ ConnID: hs.ConnID, MaxReceiveBytes: hs.MaxReceiveBytes, @@ -50,11 +83,18 @@ func (u *appUplink) OnHandshakeComplete(ctx context.Context, hs port.HandshakeIn } func (u *appUplink) OnDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason) { + lk := u.epLife(conn.EndpointID) + lk.Lock() + defer lk.Unlock() if u.presence != nil { u.presence.ClearWatch(conn.ConnID) } + u.hsMu.Lock() + _, handshook := u.handshake[conn.ConnID] + delete(u.handshake, conn.ConnID) + u.hsMu.Unlock() live, ok := u.conns.Current(conn.EndpointID) - isCurrent := ok && live.ConnID == conn.ConnID + isCurrent := ok && live.ConnID == conn.ConnID && handshook if err := u.msg.OnDisconnect(ctx, conn.EndpointID, conn.ConnID, isCurrent); err != nil { u.log.Error("message disconnect", "endpoint", conn.EndpointID, "err", err) } @@ -280,7 +320,7 @@ func (u *appUplink) publishResp(ctx context.Context, conn port.ConnInfo, resp pr return } if live, ok := u.conns.Current(conn.EndpointID); ok && live.ConnID == conn.ConnID { - limit := respPayloadLimit(live.MaxPacketSize, live.MaxReceiveBytes) + limit := broker.EffectivePayloadLimit(live.MaxPacketSize, live.MaxReceiveBytes) if limit > 0 && len(b) > limit { tooLarge := protocol.Resp{ V: protocol.Version, @@ -300,20 +340,6 @@ func (u *appUplink) publishResp(ctx context.Context, conn port.ConnInfo, resp pr } } -func respPayloadLimit(maxPacketSize uint32, maxRecvBytes int) int { - limit := 0 - if maxRecvBytes > 0 { - limit = maxRecvBytes - } - if maxPacketSize > 0 { - n := int(maxPacketSize) - if limit == 0 || n < limit { - limit = n - } - } - return limit -} - func peekRID(payload []byte) string { var peek struct { RID string `json:"rid"` diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index bbea34b..18e647d 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1207,3 +1207,75 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 - 原因:默认 info 下第二个 CONNECT、3.1.1 发到错误主题等会把整包写入 JSON 日志。 - 备选方案:改 mochi 日志调用点(需 fork)。 - 影响:排障时看不到载荷与密码,只见摘要。 + +### 复审修复 B-03 + +- 日期:2026-09-30 +- 原条款:DEVELOPMENT 第 5 节每端串行队列;Gitea #10。不改 `PublishDown` 签名。 +- 实际做法:每连接独立下行队列(256 帧 / 16MiB)和发送 goroutine。`PublishDown` 只入队;发送与上行读循环解耦。队列满返回 `ErrBackpressure`。 +- 原因:同连接同步 `InjectPacket` 与读循环写 PUBACK 会互相等待。 +- 备选方案:改 `PublishDown` 签名或继续用 20ms sleep。 +- 影响:调用方入队即返回;慢客户端只挡住该连接的发送 goroutine。 + +### 复审修复 B-06 + +- 日期:2026-09-30 +- 原条款:Gitea #13。`PublishDown` 校验当前连接与下行订阅;导出有效载荷上限。 +- 实际做法:非空 `connID` 必须仍是当前连接。未订阅 down 返回 `ErrNotSubscribed`。导出 `EffectivePayloadLimit`(Maximum Packet Size 减 128 字节包头预留)。`uplink.publishResp` 改用该函数。新连接建立时把旧连接标为 `superseded`。 +- 原因:旧连接或未订阅时写入会静默失败或写错连接。 +- 备选方案:发送时再检查(入队后连接可能已换)。 +- 影响:无订阅时下行立即失败,不再占用大帧名额。 + +### 复审修复 B-04 + +- 日期:2026-09-30 +- 原条款:Gitea #11。写出后再断开,不用固定 sleep。`serve.go` 只改 `identity.New` 的 ConnControl。 +- 实际做法:`PublishThenDisconnect` 把帧与断开原因一并入队,发送 goroutine 写完再 `Disconnect`。logout / fatalKick 改走该原语。`identity.New` 的 `ConnControl` 置 nil,踢线仍走 Session 钩子。 +- 原因:固定 20ms/50ms sleep 在慢客户端上会先断开,在快路径上又多余等待。 +- 备选方案:继续 sleep;或改 identity 生命周期(本线不改)。 +- 影响:identity 未接 ConnControl 时不再自己 20ms 踢线,生产路径统一由 Session 写出后断开。 + +### 复审修复 B-09 + +- 日期:2026-09-30 +- 原条款:Gitea #16。生命周期串行化。不改 presence/app.go。 +- 实际做法:broker 与 `appUplink` 按端编号加锁串行 `OnSessionEstablished` / `OnDisconnect` / 握手。uplink 另记 hello 握手表,仅已握手连接的断开才按当前连接通知消息线。 +- 原因:顶号时旧连接 `OnDisconnect` 可能和新连接登记交错。 +- 备选方案:改 presence 在线表(超出本线允许文件)。 +- 影响:未 hello 的断开不再把消息连接表当成已握手在线来清推送标记。 + +### 复审修复 B-10 + +- 日期:2026-09-30 +- 原条款:Gitea #17。登录写库条件更新;hello 重读令牌。不改 identity/self.go。 +- 实际做法:密码登录 `UPDATE ... WHERE COALESCE(session_hash,'') = 读到的旧值`,影响行数为 0 则 `ErrSessionWriteConflict`。hello 用 `TokenMatchesDB` 核对明文,库已被换则响应里不带回旧令牌。 +- 原因:两处同时密码登录会互相覆盖;hello 可能把已作废明文交给客户端。 +- 备选方案:写库后无条件返回本次签发明文。 +- 影响:写冲突时 OnConnect 返回 error(不回 0x86),客户端按网络故障重连。 + +### 复审修复 B-11 + +- 日期:2026-09-30 +- 原条款:Gitea #18。令牌闲置按在线计。 +- 实际做法:闲置判断取 `session_used_at` / `online_since` / `offline_since` 的较新者;当前在线(`online_since >= offline_since`)视为未闲置。 +- 原因:只看 `session_used_at` 会让长期在线却很少写库的令牌过期。 +- 备选方案:在线时每小时强制刷新 used_at(已有 touch,但仍可能窗口不够)。 +- 影响:在线设备不会因为闲置天数被踢;离线后从最后一次在线/离线时刻起算。 + +### 复审修复 B-12 + +- 日期:2026-09-30 +- 原条款:Gitea #19。认证超时与每端校验并发。不改 auth 池/PHC。 +- 实际做法:`Authenticate` 套 30 秒超时;argon2 `Verify` 前每端信号量 2。`OnConnect` 同样带 30 秒 ctx。 +- 原因:慢哈希或卡住的校验会堵住 mochi 读循环;同一编号并发登录会打满全局哈希池。 +- 备选方案:改全局 Pool 大小(超出允许文件)。 +- 影响:超时表现为内部错误断开(不回 0x86)。 + +### 复审修复 B-08 + +- 日期:2026-09-30 +- 原条款:Gitea #15。Shutdown API。完整 HTTP 停机依赖 L-03。 +- 实际做法:`Broker.Shutdown` 对现有连接发 MQTT 5 `0x8B`,清空上行队列并 `Close`。`serve` 在 listener Close 之前调用。HTTP `Shutdown` 留给 L-03。 +- 原因:只关 listener 时 MQTT 客户端看不到规范的停机原因码。 +- 备选方案:等 L-03 一并做(本线仍提供 broker API,避免监听线无法调用)。 +- 影响:进程退出时端会收到 server shutting down;监听器 HTTP 优雅停机仍未做。 diff --git a/internal/broker/authn.go b/internal/broker/authn.go index 29de669..265643f 100644 --- a/internal/broker/authn.go +++ b/internal/broker/authn.go @@ -32,6 +32,9 @@ type Login struct { // 内存中的 session_used_at(毫秒)与上次落库时间。 usedAt map[string]int64 lastFlush map[string]int64 + + verifyMu sync.Mutex + verifySem map[string]chan struct{} } // LoginOptions 装配 Login。 @@ -67,9 +70,15 @@ func NewLogin(opts LoginOptions) *Login { Now: now, usedAt: make(map[string]int64), lastFlush: make(map[string]int64), + verifySem: make(map[string]chan struct{}), } } +const ( + authTimeout = 30 * time.Second + verifyPerEndpoint = 2 +) + // Authenticate 按 DEVELOPMENT 第 5 节校验;内部故障返回 error。 func (l *Login) Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error) { if l == nil || l.DB == nil { @@ -78,6 +87,11 @@ func (l *Login) Authenticate(ctx context.Context, endpointID string, password [] if endpointID == "" { return AuthResult{OK: false}, nil } + if ctx == nil { + ctx = context.Background() + } + ctx, cancel := context.WithTimeout(ctx, authTimeout) + defer cancel() row, err := l.loadEndpoint(ctx, endpointID) if err != nil { @@ -107,6 +121,8 @@ type endpointAuthRow struct { loginHash string sessionHash []byte // 原始 32 字节;无令牌时 nil sessionUsedAt int64 // 毫秒;无则 0 + onlineSince int64 + offlineSince int64 } func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, error) { @@ -115,10 +131,12 @@ func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, e enabled int sessHex sql.NullString usedAt sql.NullInt64 + online sql.NullInt64 + offline sql.NullInt64 ) err := l.DB.Read.QueryRowContext(ctx, ` -SELECT login_hash, enabled, session_hash, session_used_at -FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt) +SELECT login_hash, enabled, session_hash, session_used_at, online_since, offline_since +FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt, &online, &offline) if err != nil { if errors.Is(err, sql.ErrNoRows) { return endpointAuthRow{}, ErrEndpointNotFound @@ -132,6 +150,12 @@ FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt) if usedAt.Valid { row.sessionUsedAt = usedAt.Int64 } + if online.Valid { + row.onlineSince = online.Int64 + } + if offline.Valid { + row.offlineSince = offline.Int64 + } if sessHex.Valid && sessHex.String != "" { raw, decErr := hex.DecodeString(sessHex.String) if decErr != nil || len(raw) != 32 { @@ -160,11 +184,8 @@ func (l *Login) authSession(ctx context.Context, endpointID, token string, row e usedAt = mem } l.usedMu.Unlock() - if l.IdleDays > 0 { - idle := time.Duration(l.IdleDays) * 24 * time.Hour - if usedAt <= 0 || now.Sub(time.UnixMilli(usedAt)) > idle { - return false, nil - } + if l.IdleDays > 0 && !sessionIdleOK(now, usedAt, row.onlineSince, row.offlineSince, l.IdleDays) { + return false, nil } if err := l.touchSessionUsed(ctx, endpointID, nowMs); err != nil { return false, err @@ -204,7 +225,11 @@ func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP if l.Pool == nil { return false, "", errors.New("broker: password pool not configured") } + if err := l.acquireVerify(ctx, endpointID); err != nil { + return false, "", err + } match, verErr := l.Pool.Verify(ctx, auth.PasswordLogin, password, row.loginHash) + l.releaseVerify(endpointID) if verErr != nil { return false, "", verErr } @@ -220,12 +245,28 @@ func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP } nowMs := l.Now().UnixMilli() hashHex := hex.EncodeToString(hash) + var oldHex any + if len(row.sessionHash) == 0 { + oldHex = "" + } else { + oldHex = hex.EncodeToString(row.sessionHash) + } writeErr := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error { - _, e := tx.Exec(` + res, e := tx.Exec(` UPDATE endpoints SET session_hash = ?, session_issued_at = ?, session_used_at = ? -WHERE id = ?`, hashHex, nowMs, nowMs, endpointID) - return e +WHERE id = ? AND COALESCE(session_hash, '') = ?`, hashHex, nowMs, nowMs, endpointID, oldHex) + if e != nil { + return e + } + n, nErr := res.RowsAffected() + if nErr != nil { + return nErr + } + if n == 0 { + return ErrSessionWriteConflict + } + return nil }) if writeErr != nil { return false, "", writeErr @@ -288,6 +329,73 @@ func (l *Login) SessionHashOf(ctx context.Context, endpointID string) ([]byte, e return hex.DecodeString(sessHex.String) } +// TokenMatchesDB 握手时重读:明文令牌是否仍对应库中当前哈希。 +func (l *Login) TokenMatchesDB(ctx context.Context, endpointID, token string) (bool, error) { + if l == nil || l.DB == nil || token == "" { + return false, nil + } + got := l.Tokens.HashToken(token) + dbHash, err := l.SessionHashOf(ctx, endpointID) + if err != nil { + return false, err + } + if len(dbHash) == 0 { + return false, nil + } + return auth.EqualHash(got, dbHash), nil +} + +func sessionIdleOK(now time.Time, usedAt, onlineSince, offlineSince int64, idleDays int) bool { + if idleDays <= 0 { + return true + } + online := onlineSince > 0 && onlineSince >= offlineSince + if online { + return true + } + activity := usedAt + if onlineSince > activity { + activity = onlineSince + } + if offlineSince > activity { + activity = offlineSince + } + if activity <= 0 { + return false + } + idle := time.Duration(idleDays) * 24 * time.Hour + return now.Sub(time.UnixMilli(activity)) <= idle +} + +func (l *Login) acquireVerify(ctx context.Context, endpointID string) error { + l.verifyMu.Lock() + sem := l.verifySem[endpointID] + if sem == nil { + sem = make(chan struct{}, verifyPerEndpoint) + l.verifySem[endpointID] = sem + } + l.verifyMu.Unlock() + select { + case sem <- struct{}{}: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func (l *Login) releaseVerify(endpointID string) { + l.verifyMu.Lock() + sem := l.verifySem[endpointID] + l.verifyMu.Unlock() + if sem == nil { + return + } + select { + case <-sem: + default: + } +} + // LooksLikeSessionToken 暴露给测试。 func (l *Login) LooksLikeSessionToken(s string) bool { return strings.HasPrefix(s, "nst_") diff --git a/internal/broker/authn_idle_test.go b/internal/broker/authn_idle_test.go new file mode 100644 index 0000000..46c7b9e --- /dev/null +++ b/internal/broker/authn_idle_test.go @@ -0,0 +1,123 @@ +package broker + +import ( + "context" + "database/sql" + "sync" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/store" +) + +func TestSessionIdleOKUsesOnlineOffline(t *testing.T) { + now := time.UnixMilli(1_700_000_000_000) + idleDays := 1 + old := now.Add(-48 * time.Hour).UnixMilli() + recentOffline := now.Add(-2 * time.Hour).UnixMilli() + if sessionIdleOK(now, old, 0, 0, idleDays) { + t.Fatal("stale used_at should expire") + } + if !sessionIdleOK(now, old, now.UnixMilli(), 0, idleDays) { + t.Fatal("currently online should not expire") + } + if !sessionIdleOK(now, old, 0, recentOffline, idleDays) { + t.Fatal("recent offline_since should keep token") + } +} + +func TestPasswordLoginConditionalUpdate(t *testing.T) { + dir := t.TempDir() + db, err := store.Open(dir, "FULL") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + pool := auth.NewStubHashPool() + login := NewLogin(LoginOptions{DB: db, Pool: pool, Tokens: auth.NewSessionTokens(), Locks: auth.NewLoginLocks(), IdleDays: 30}) + phc, _ := pool.Hash(context.Background(), auth.PasswordLogin, "password1") + _ = db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + _, e := tx.Exec(`INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at) +VALUES ('ep-cond', '', ?, NULL, 0, 0, 1, ?)`, phc, time.Now().UnixMilli()) + return e + }) + res, err := login.Authenticate(context.Background(), "ep-cond", []byte("password1"), "1.1.1.1") + if err != nil || !res.OK || res.SessionToken == "" { + t.Fatalf("first login %+v err=%v", res, err) + } + ok, err := login.TokenMatchesDB(context.Background(), "ep-cond", res.SessionToken) + if err != nil || !ok { + t.Fatalf("match=%v err=%v", ok, err) + } + res2, err := login.Authenticate(context.Background(), "ep-cond", []byte("password1"), "1.1.1.1") + if err != nil || !res2.OK { + t.Fatalf("second login %+v err=%v", res2, err) + } + ok, _ = login.TokenMatchesDB(context.Background(), "ep-cond", res.SessionToken) + if ok { + t.Fatal("old token should not match after second login") + } + ok, _ = login.TokenMatchesDB(context.Background(), "ep-cond", res2.SessionToken) + if !ok { + t.Fatal("new token should match") + } +} + +func TestAuthenticateRespectsCanceledContext(t *testing.T) { + dir := t.TempDir() + db, err := store.Open(dir, "FULL") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + pool := &blockingPool{ready: make(chan struct{}), release: make(chan struct{})} + login := NewLogin(LoginOptions{DB: db, Pool: pool, Tokens: auth.NewSessionTokens(), Locks: auth.NewLoginLocks()}) + phc, _ := auth.NewStubHashPool().Hash(context.Background(), auth.PasswordLogin, "password1") + _ = db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + _, e := tx.Exec(`INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at) +VALUES ('ep-to', '', ?, NULL, 0, 0, 1, ?)`, phc, time.Now().UnixMilli()) + return e + }) + ctx, cancel := context.WithCancel(context.Background()) + var wg sync.WaitGroup + wg.Add(1) + var gotErr error + go func() { + defer wg.Done() + _, gotErr = login.Authenticate(ctx, "ep-to", []byte("password1"), "9.9.9.9") + }() + select { + case <-pool.ready: + case <-time.After(2 * time.Second): + t.Fatal("verify did not start") + } + cancel() + wg.Wait() + close(pool.release) + if gotErr == nil { + t.Fatal("expected canceled auth") + } +} + +type blockingPool struct { + ready chan struct{} + release chan struct{} + once sync.Once +} + +func (p *blockingPool) Hash(ctx context.Context, kind auth.PasswordKind, password string) (string, error) { + return auth.NewStubHashPool().Hash(ctx, kind, password) +} + +func (p *blockingPool) Verify(ctx context.Context, _ auth.PasswordKind, _, _ string) (bool, error) { + p.once.Do(func() { close(p.ready) }) + select { + case <-ctx.Done(): + return false, ctx.Err() + case <-p.release: + return true, nil + } +} + +func (p *blockingPool) QueueLen() int { return 0 } diff --git a/internal/broker/b02_large_test.go b/internal/broker/b02_large_test.go index f353b73..45f1655 100644 --- a/internal/broker/b02_large_test.go +++ b/internal/broker/b02_large_test.go @@ -3,6 +3,7 @@ package broker import ( "bytes" "context" + "errors" "io" "net" "sync" @@ -81,7 +82,16 @@ func TestLargeFrameQuotaReleasedOnDisconnect(t *testing.T) { t.Fatal(err) } cancel() - if len(b.largeSem) == 0 { + held := false + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + if len(b.largeSem) > 0 { + held = true + break + } + time.Sleep(10 * time.Millisecond) + } + if !held { t.Fatal("expected a held large slot before disconnect") } _ = w.Close() @@ -90,7 +100,7 @@ func TestLargeFrameQuotaReleasedOnDisconnect(t *testing.T) { case <-time.After(3 * time.Second): t.Fatal("client attach did not return") } - deadline := time.Now().Add(2 * time.Second) + deadline = time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { if len(b.largeSem) == 0 { return @@ -119,11 +129,11 @@ func TestLargeFrameQuotaReleasedWithoutSubscriber(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), time.Second) err := b.PublishDown(ctx, "ep-large-nosub", "", payload, port.PublishOpts{QoS: 1}) cancel() - if err != nil { - t.Fatal(err) + if !errors.Is(err, ErrNotSubscribed) { + t.Fatalf("err=%v want ErrNotSubscribed", err) } - if n := len(b.largeSem); n != 0 { - t.Fatalf("held slots without subscriber: %d", n) + if len(b.largeSem) != 0 { + t.Fatalf("held slots without subscriber: %d", len(b.largeSem)) } } diff --git a/internal/broker/b03_downlink_test.go b/internal/broker/b03_downlink_test.go new file mode 100644 index 0000000..577084a --- /dev/null +++ b/internal/broker/b03_downlink_test.go @@ -0,0 +1,212 @@ +package broker + +import ( + "bytes" + "context" + "errors" + "io" + "net" + "strconv" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" + "github.com/mochi-mqtt/server/v2/packets" +) + +func TestPublishDownBackpressureWhenQueueFull(t *testing.T) { + // 无缓冲 pipe:发送 goroutine 在客户端不读时堵住,队列才能填满。 + b, err := New(Options{Authenticator: AllowAuthenticator{}}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b.Close() }() + r, w := net.Pipe() + done := make(chan struct{}) + go func() { + defer close(done) + _ = b.AttachTCP(r) + }() + defer func() { + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + }() + connectAndSubscribe(t, w, "ep-bp", 0) + waitSession(t, b, "ep-bp") + + payload := bytes.Repeat([]byte("q"), 1024) + var sawBP bool + start := time.Now() + // mochi outbound 缓冲 1024,发送 goroutine 要先填满它才会堵住,随后才轮到本地下行队列。 + for i := 0; i < 1024+downQueueMax+16; i++ { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + err := b.PublishDown(ctx, "ep-bp", "", payload, port.PublishOpts{QoS: 1}) + cancel() + if errors.Is(err, ErrBackpressure) { + if time.Since(start) > 50*time.Millisecond && i == 0 { + t.Fatalf("first backpressure took %s", time.Since(start)) + } + sawBP = true + break + } + if err != nil { + t.Fatalf("publish %d: %v", i, err) + } + } + if !sawBP { + t.Fatal("expected ErrBackpressure") + } +} + +func TestSlowClientDoesNotBlockOtherPublishDown(t *testing.T) { + b, err := New(Options{Authenticator: AllowAuthenticator{}}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = b.Close() }() + + slowW, slowDone := acceptAndDial(t, b) + fastW, fastDone := acceptAndDial(t, b) + defer func() { + _ = slowW.Close() + _ = fastW.Close() + select { + case <-slowDone: + case <-time.After(3 * time.Second): + } + select { + case <-fastDone: + case <-time.After(3 * time.Second): + } + }() + + writeConnect(t, slowW, "ep-slow", 30, 0) + readExactPacket(t, slowW, packets.Connack, 3*time.Second) + writeSubscribe(t, slowW, downTopic("ep-slow")) + readExactPacket(t, slowW, packets.Suback, 3*time.Second) + writeConnect(t, fastW, "ep-fast", 30, 0) + readExactPacket(t, fastW, packets.Connack, 3*time.Second) + writeSubscribe(t, fastW, downTopic("ep-fast")) + readExactPacket(t, fastW, packets.Suback, 3*time.Second) + waitSession(t, b, "ep-slow") + waitSession(t, b, "ep-fast") + + big := bytes.Repeat([]byte("s"), 64*1024) + for i := 0; i < 8; i++ { + _ = b.PublishDown(context.Background(), "ep-slow", "", big, port.PublishOpts{QoS: 1}) + } + + small := []byte(`{"v":1,"type":"resp"}`) + start := time.Now() + if err := b.PublishDown(context.Background(), "ep-fast", "", small, port.PublishOpts{QoS: 1}); err != nil { + t.Fatalf("fast publish: %v", err) + } + if time.Since(start) > 100*time.Millisecond { + t.Fatalf("fast PublishDown took %s", time.Since(start)) + } + got := readDownPayload(t, fastW, 2*time.Second) + if !bytes.Equal(got, small) { + t.Fatalf("fast got %q", got) + } +} + +func TestDownlinkFIFOOrder(t *testing.T) { + b, w, done := startTCPClient(t, "ep-ord") + defer func() { _ = b.Close() }() + defer func() { + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + }() + writeConnect(t, w, "ep-ord", 30, 0) + readExactPacket(t, w, packets.Connack, 3*time.Second) + writeSubscribe(t, w, downTopic("ep-ord")) + readExactPacket(t, w, packets.Suback, 3*time.Second) + waitSession(t, b, "ep-ord") + + const n = 64 + for i := 0; i < n; i++ { + p := []byte(strconv.Itoa(i)) + if err := b.PublishDown(context.Background(), "ep-ord", "", p, port.PublishOpts{QoS: 1}); err != nil { + t.Fatal(err) + } + } + for i := 0; i < n; i++ { + got := readDownPayload(t, w, 3*time.Second) + want := []byte(strconv.Itoa(i)) + if !bytes.Equal(got, want) { + t.Fatalf("order %d: got %s want %s", i, got, want) + } + } +} + +func acceptAndDial(t *testing.T, b *Broker) (net.Conn, chan struct{}) { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = ln.Close() }) + done := make(chan struct{}) + go func() { + defer close(done) + c, accErr := ln.Accept() + if accErr != nil { + return + } + _ = b.AttachTCP(c) + }() + w, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + t.Fatal(err) + } + return w, done +} + +func readDownPayload(t *testing.T, conn net.Conn, timeout time.Duration) []byte { + t.Helper() + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + _ = conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond)) + hdr := make([]byte, 1) + if _, err := io.ReadFull(conn, hdr); err != nil { + continue + } + rem, err := readRemainingLengthConn(conn) + if err != nil { + continue + } + body := make([]byte, rem) + if _, err := io.ReadFull(conn, body); err != nil { + continue + } + if hdr[0]>>4 != packets.Publish { + continue + } + qos := (hdr[0] >> 1) & 0x3 + pk := new(packets.Packet) + pk.ProtocolVersion = 5 + pk.FixedHeader = packets.FixedHeader{Type: packets.Publish, Remaining: rem, Qos: qos} + if err := pk.PublishDecode(body); err != nil { + t.Fatal(err) + } + if qos > 0 { + ack := packets.Packet{ + FixedHeader: packets.FixedHeader{Type: packets.Puback}, + ProtocolVersion: 5, + PacketID: pk.PacketID, + } + var ab bytes.Buffer + _ = ack.PubackEncode(&ab) + _, _ = conn.Write(ab.Bytes()) + } + return pk.Payload + } + t.Fatal("timeout waiting publish") + return nil +} diff --git a/internal/broker/b04_b06_b08_test.go b/internal/broker/b04_b06_b08_test.go new file mode 100644 index 0000000..806f7e7 --- /dev/null +++ b/internal/broker/b04_b06_b08_test.go @@ -0,0 +1,128 @@ +package broker + +import ( + "bytes" + "context" + "errors" + "io" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" + "github.com/mochi-mqtt/server/v2/packets" +) + +func TestPublishDownRequiresDownSubscription(t *testing.T) { + b, w, done := startTCPClient(t, "ep-nosub") + defer func() { _ = b.Close() }() + defer func() { + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + }() + writeConnect(t, w, "ep-nosub", 30, 0) + readExactPacket(t, w, packets.Connack, 3*time.Second) + waitSession(t, b, "ep-nosub") + err := b.PublishDown(context.Background(), "ep-nosub", "", []byte(`{"v":1}`), port.PublishOpts{QoS: 0}) + if !errors.Is(err, ErrNotSubscribed) { + t.Fatalf("err=%v want ErrNotSubscribed", err) + } +} + +func TestPublishDownRejectsStaleConnID(t *testing.T) { + b, w, done := startTCPClient(t, "ep-stale") + defer func() { _ = b.Close() }() + defer func() { + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + }() + writeConnect(t, w, "ep-stale", 30, 0) + readExactPacket(t, w, packets.Connack, 3*time.Second) + writeSubscribe(t, w, downTopic("ep-stale")) + readExactPacket(t, w, packets.Suback, 3*time.Second) + waitSession(t, b, "ep-stale") + err := b.PublishDown(context.Background(), "ep-stale", "dead-conn", []byte(`{"v":1}`), port.PublishOpts{QoS: 0}) + if !errors.Is(err, ErrNoConnection) { + t.Fatalf("err=%v want ErrNoConnection", err) + } +} + +func TestPublishThenDisconnectWritesThenCloses(t *testing.T) { + b, w, done := startTCPClient(t, "ep-ptd") + defer func() { _ = b.Close() }() + defer func() { + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + }() + writeConnect(t, w, "ep-ptd", 30, 0) + readExactPacket(t, w, packets.Connack, 3*time.Second) + writeSubscribe(t, w, downTopic("ep-ptd")) + readExactPacket(t, w, packets.Suback, 3*time.Second) + waitSession(t, b, "ep-ptd") + + payload := []byte(`{"v":1,"type":"fatal","reason":"disabled"}`) + if err := b.PublishThenDisconnect(context.Background(), "ep-ptd", "", payload, 1, port.DisconnectFatal); err != nil { + t.Fatal(err) + } + got := readDownPayload(t, w, 3*time.Second) + if !bytes.Equal(got, payload) { + t.Fatalf("got %s", got) + } + _ = 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() { + _ = w.Close() + select { + case <-done: + case <-time.After(3 * time.Second): + } + }() + writeConnect(t, w, "ep-shut", 30, 0) + readExactPacket(t, w, packets.Connack, 3*time.Second) + waitSession(t, b, "ep-shut") + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + if err := b.Shutdown(ctx); err != nil { + t.Fatal(err) + } + _ = w.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 32) + n, err := io.ReadAtLeast(w, buf, 2) + if err != nil && n == 0 { + return + } + if n > 0 && buf[0]>>4 != packets.Disconnect { + t.Fatalf("want disconnect got %x", buf[:n]) + } +} + +func TestEffectivePayloadLimitSubtractsOverhead(t *testing.T) { + got := EffectivePayloadLimit(200, 0) + if got != 200-packetOverheadBudget { + t.Fatalf("got %d", got) + } + got = EffectivePayloadLimit(200, 50) + if got != 50 { + t.Fatalf("got %d want 50", got) + } +} diff --git a/internal/broker/broker.go b/internal/broker/broker.go index 91accc6..1474d69 100644 --- a/internal/broker/broker.go +++ b/internal/broker/broker.go @@ -37,7 +37,20 @@ var ErrNoConnection = errors.New("broker: no active connection") // ErrLargeFrameTimeout 全局大帧名额在有界等待内拿不到。 var ErrLargeFrameTimeout = errors.New("broker: large frame quota timeout") -const largeAcquireWait = 5 * time.Second +// ErrBackpressure 该连接下行队列已满(帧数或字节数)。 +var ErrBackpressure = errors.New("broker: downlink backpressure") + +// ErrNotSubscribed 当前连接尚未订阅下行主题。 +var ErrNotSubscribed = errors.New("broker: down topic not subscribed") + +// ErrSessionWriteConflict 密码登录写令牌时发现库已被并发更新。 +var ErrSessionWriteConflict = errors.New("broker: session token write conflict") + +const ( + largeAcquireWait = 5 * time.Second + downQueueMax = 256 + downQueueBytes = 16 << 20 +) // AuthResult 是登录校验结论(N3 实现真实逻辑;N2 默认拒绝)。 type AuthResult struct { @@ -100,6 +113,9 @@ type Broker struct { largeSem chan struct{} closed atomic.Bool + + lifeMu sync.Mutex + lifeLocks map[string]*sync.Mutex } type connState struct { @@ -120,6 +136,13 @@ type connState struct { metricsCounted bool established bool createdAt time.Time + closing bool + superseded bool + downCh chan downItem + downStop chan struct{} + downDone chan struct{} + downBytes atomic.Int64 + sentPub atomic.Int64 mu sync.Mutex handshakeTimer *time.Timer @@ -162,18 +185,19 @@ func New(opts Options) (*Broker, error) { }) b := &Broker{ - server: srv, - auth: auth, - uplink: uplink, - log: log, - onDrop: opts.OnPublishDropped, - metrics: opts.Metrics, - current: make(map[string]*connState), - byClient: make(map[*mqtt.Client]*connState), - byConnID: make(map[port.ConnID]*connState), - closedCh: make(chan struct{}), - queues: make(map[string]*uplinkQueue), - largeSem: make(chan struct{}, largeFrameSlots), + server: srv, + auth: auth, + uplink: uplink, + log: log, + onDrop: opts.OnPublishDropped, + metrics: opts.Metrics, + current: make(map[string]*connState), + byClient: make(map[*mqtt.Client]*connState), + byConnID: make(map[port.ConnID]*connState), + closedCh: make(chan struct{}), + queues: make(map[string]*uplinkQueue), + largeSem: make(chan struct{}, largeFrameSlots), + lifeLocks: make(map[string]*sync.Mutex), } b.hook = &nixHook{b: b} if err := srv.AddHook(b.hook, nil); err != nil { @@ -203,10 +227,36 @@ func (b *Broker) Close() error { for _, q := range b.queues { q.close() } + b.queues = make(map[string]*uplinkQueue) b.queuesMu.Unlock() return b.server.Close() } +// Shutdown 向所有连接发 MQTT 5 0x8B 后关闭。完整 HTTP 停机顺序见 L-03。 +func (b *Broker) Shutdown(ctx context.Context) error { + if b.closed.Load() { + return nil + } + b.connsMu.RLock() + clients := make([]*mqtt.Client, 0, len(b.byClient)) + for cl := range b.byClient { + if cl != nil { + clients = append(clients, cl) + } + } + b.connsMu.RUnlock() + for _, cl := range clients { + _ = b.server.DisconnectClient(cl, packets.ErrServerShuttingDown) + } + if ctx != nil { + select { + case <-ctx.Done(): + default: + } + } + return b.Close() +} + // AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。 func (b *Broker) AttachTCP(conn net.Conn) error { return b.server.EstablishConnection("tcp", conn) @@ -219,44 +269,55 @@ func (b *Broker) AttachWS(conn net.Conn) error { // PublishDown 实现 port.Downlink。 func (b *Broker) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error { - if b.closed.Load() { - return errors.New("broker: closed") - } - st := b.lookupConn(endpointID, connID) - if st == nil { - return ErrNoConnection - } - - limit := effectivePayloadLimit(st.maxPacketSize, st.maxRecvBytes) - if limit > 0 && len(payload) > limit { - return ErrPayloadTooLarge - } - qos := opts.QoS if qos > 1 { qos = 1 } - topic := downTopic(endpointID) - large := len(payload) > largeFrameBytes - if large { - if err := b.acquireLarge(ctx); err != nil { - return err - } - st.mu.Lock() - st.largePending++ - st.mu.Unlock() + return b.enqueueDownlink(endpointID, connID, payload, qos, "") +} + +// PublishThenDisconnect 把一帧写入该连接下行队列,写出后再断开(无固定 sleep)。 +func (b *Broker) PublishThenDisconnect(_ context.Context, endpointID string, connID port.ConnID, payload []byte, qos byte, reason port.DisconnectReason) error { + if qos > 1 { + qos = 1 + } + if reason == "" { + reason = port.DisconnectNormal + } + return b.enqueueDownlink(endpointID, connID, payload, qos, reason) +} + +func (b *Broker) enqueueDownlink(endpointID string, connID port.ConnID, payload []byte, qos byte, disconnect port.DisconnectReason) error { + if b.closed.Load() { + return errors.New("broker: closed") + } + st := b.lookupCurrent(endpointID, connID) + if st == nil { + return ErrNoConnection } - if err := b.server.Publish(topic, payload, false, qos); err != nil { - if large { - b.finishLargePublish(st) - } - return err + st.mu.Lock() + maxRecv := st.maxRecvBytes + closing := st.closing + superseded := st.superseded + st.mu.Unlock() + if closing || superseded { + return ErrNoConnection } - if large { - b.finishLargePublish(st) + if !b.hasDownSub(st) { + return ErrNotSubscribed } - return nil + + limit := EffectivePayloadLimit(st.maxPacketSize, maxRecv) + if limit > 0 && len(payload) > limit { + return ErrPayloadTooLarge + } + + return st.enqueueDown(downItem{ + payload: append([]byte(nil), payload...), + qos: qos, + disconnect: disconnect, + }) } func (b *Broker) acquireLarge(ctx context.Context) error { @@ -367,6 +428,31 @@ func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState { return b.current[endpointID] } +// lookupCurrent 只返回该端当前连接;connID 非空时必须仍是当前连接。 +func (b *Broker) lookupCurrent(endpointID string, connID port.ConnID) *connState { + b.connsMu.RLock() + defer b.connsMu.RUnlock() + cur := b.current[endpointID] + if cur == nil { + return nil + } + if connID != "" && cur.connID != connID { + return nil + } + return cur +} + +func (b *Broker) endpointLife(endpointID string) *sync.Mutex { + b.lifeMu.Lock() + defer b.lifeMu.Unlock() + m := b.lifeLocks[endpointID] + if m == nil { + m = &sync.Mutex{} + b.lifeLocks[endpointID] = m + } + return m +} + func downTopic(endpointID string) string { return "nix/c/" + endpointID + "/down" } @@ -375,6 +461,11 @@ func upTopic(endpointID string) string { return "nix/c/" + endpointID + "/up" } +// EffectivePayloadLimit 下行载荷上限:客户端 Maximum Packet Size 减包头预留,再与 max_receive_bytes 取更严者。 +func EffectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int { + return effectivePayloadLimit(maxPacketSize, maxRecvBytes) +} + func effectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int { limit := 0 if maxPacketSize > 0 { diff --git a/internal/broker/downlink.go b/internal/broker/downlink.go new file mode 100644 index 0000000..5d31ca6 --- /dev/null +++ b/internal/broker/downlink.go @@ -0,0 +1,192 @@ +package broker + +import ( + "context" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/app/port" +) + +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) + before := st.sentPub.Load() + var err error + for { + if b.closed.Load() { + break + } + 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 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.waitPacketWritten(before) + _ = 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) signalSent(item downItem) { + if item.sent == nil { + return + } + select { + case <-item.sent: + default: + close(item.sent) + } +} diff --git a/internal/broker/hooks.go b/internal/broker/hooks.go index b6795fb..567a31b 100644 --- a/internal/broker/hooks.go +++ b/internal/broker/hooks.go @@ -30,6 +30,7 @@ func (h *nixHook) Provides(b byte) bool { mqtt.OnQosComplete, mqtt.OnQosDropped, mqtt.OnSubscribed, + mqtt.OnPacketSent, }, []byte{b}) } @@ -79,7 +80,9 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { "endpoint", endpointID, "receive_maximum", rm) } - res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP) + authCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + res, err := h.b.auth.Authenticate(authCtx, endpointID, pk.Connect.Password, remoteIP) if err != nil { return err // mochi 不回 CONNACK,直接断开;不登记连接表 } @@ -171,6 +174,18 @@ 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) + } +} + func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) { h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload)) h.b.connsMu.RLock() @@ -188,7 +203,9 @@ func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) { func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) { h.b.connsMu.Lock() st := h.b.byClient[cl] + var old *connState if st != nil { + old = h.b.current[st.endpointID] h.b.current[st.endpointID] = st st.established = true } @@ -196,6 +213,15 @@ func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) { if st == nil { return } + lk := h.b.endpointLife(st.endpointID) + lk.Lock() + if old != nil && old != st { + old.mu.Lock() + old.superseded = true + old.mu.Unlock() + } + lk.Unlock() + st.startDownLoop(h.b) info := port.ConnInfo{ ConnID: st.connID, EndpointID: st.endpointID, @@ -224,6 +250,10 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) { if st == nil { return } + lk := h.b.endpointLife(st.endpointID) + lk.Lock() + st.stopDownLoop() + lk.Unlock() h.b.releaseAllLarge(st) h.b.cancelHandshakeDeadline(st.endpointID, st.connID) diff --git a/internal/broker/session.go b/internal/broker/session.go index 52bf739..2b19a07 100644 --- a/internal/broker/session.go +++ b/internal/broker/session.go @@ -210,6 +210,13 @@ func (s *Session) handleHello(ctx context.Context, conn port.ConnInfo, st *connS } s.b.SetMaxReceiveBytes(conn.EndpointID, conn.ConnID, maxRecv) + if conn.SessionToken != "" && s.login != nil { + if keep, chkErr := s.login.TokenMatchesDB(ctx, conn.EndpointID, conn.SessionToken); chkErr != nil { + s.log.Error("re-read session token", "endpoint", conn.EndpointID, "err", chkErr) + } else if !keep { + conn.SessionToken = "" + } + } data := protocol.HelloData{ ServerTimeMs: s.now().UnixMilli(), ServerVersion: s.limits.ServerVersion, @@ -279,14 +286,17 @@ func (s *Session) handleLogout(ctx context.Context, conn port.ConnInfo, req *pro } } resp := protocol.Resp{V: protocol.Version, Type: protocol.TypeResp, RID: req.RID, OK: true} - if err := s.publishJSON(ctx, conn, resp, 1); err != nil { - s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", err) + raw, err := protocol.Marshal(resp) + if err != nil { + s.replyErr(ctx, conn, req.RID, protocol.CodeBusy, "marshal logout resp") + return nil + } + if pubErr := s.b.PublishThenDisconnect(ctx, conn.EndpointID, conn.ConnID, raw, 1, port.DisconnectNormal); pubErr != nil { + s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", pubErr) + go func() { + _ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal) + }() } - go func() { - // 稍等让 QoS1 resp 写入连接,再断开 - time.Sleep(50 * time.Millisecond) - _ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal) - }() return nil } @@ -327,11 +337,15 @@ func (s *Session) fatalKick(ctx context.Context, endpointID, reason string) erro return nil } fatal := protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: reason} - _ = s.publishJSON(ctx, info, fatal, 1) - go func() { - time.Sleep(20 * time.Millisecond) - _ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal) - }() + raw, err := protocol.Marshal(fatal) + if err != nil { + return err + } + if pubErr := s.b.PublishThenDisconnect(ctx, info.EndpointID, info.ConnID, raw, 1, port.DisconnectFatal); pubErr != nil { + go func() { + _ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal) + }() + } return nil }