From 42160720bd49d1f456e5e3a3b472ce92318a9dad Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 15:08:03 +0800 Subject: [PATCH] =?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), + } +}