From 9a17c6a26bd7b09f875d664a48c5caf59cfcb09f Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 16:08:38 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E5=81=9C=E6=9C=BA=E5=85=88=E5=81=9C?= =?UTF-8?q?=E6=8E=A5=E5=8F=97=E5=B9=B6=E7=AD=89=E5=BE=85=E5=BE=AA=E7=8E=AF?= =?UTF-8?q?=E5=86=8D=E6=96=AD=E5=BC=80=20MQTT?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- cmd/nixmsg/serve.go | 65 ++++++++++++++++--- cmd/nixmsg/serve_shutdown_test.go | 73 +++++++++++++++++++++ cmd/nixmsg/uplink_integration_test.go | 44 ++++++++++++- deploy/docker-compose.yml | 1 + docs/DEVIATIONS.md | 9 +++ docs/OPS.md | 1 + internal/listener/server.go | 91 ++++++++++++++++++++++++--- internal/listener/shutdown_test.go | 81 ++++++++++++++++++++++++ 8 files changed, 345 insertions(+), 20 deletions(-) create mode 100644 cmd/nixmsg/serve_shutdown_test.go create mode 100644 internal/listener/shutdown_test.go diff --git a/cmd/nixmsg/serve.go b/cmd/nixmsg/serve.go index 9cd7725..6dd0985 100644 --- a/cmd/nixmsg/serve.go +++ b/cmd/nixmsg/serve.go @@ -43,7 +43,20 @@ func cmdServe(_ []string) error { setupJSONLogger(cfg.Log) ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) defer stop() - return runServe(ctx, cfg) + return runServe(withServeStop(ctx, stop), cfg) +} + +type serveStopKey struct{} + +func withServeStop(ctx context.Context, stop context.CancelFunc) context.Context { + return context.WithValue(ctx, serveStopKey{}, stop) +} + +func invokeServeStop(ctx context.Context) { + stop, _ := ctx.Value(serveStopKey{}).(context.CancelFunc) + if stop != nil { + stop() + } } func runServe(ctx context.Context, cfg config.Config) error { @@ -297,27 +310,59 @@ func runServe(ctx context.Context, cfg config.Config) error { } } - loopCtx, loopCancel := context.WithCancel(ctx) + loopCtx, loopCancel := context.WithCancel(context.Background()) defer loopCancel() - go messageLoops(loopCtx, msgApp, db, hashPool, metricsReg) + loopsDone := make(chan struct{}) + go func() { + defer close(loopsDone) + messageLoops(loopCtx, msgApp, db, hashPool, metricsReg) + }() <-ctx.Done() + invokeServeStop(ctx) + shutdownDeadline := time.Now().Add(30 * time.Second) + _ = lnSrv.StopAccept() 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() + <-loopsDone + + drainBudget := 10 * time.Second + drainStart := time.Now() + drainCtx, drainCancel := context.WithTimeout(context.Background(), drainBudget) if drainErr := db.Queue.Drain(drainCtx); drainErr != nil && !errors.Is(drainErr, context.DeadlineExceeded) { slog.Error("write queue drain", "err", drainErr) } + drainCancel() + + shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second) + _ = brk.Shutdown(shutCtx) + shutCancel() + + secondDrain := drainBudget - time.Since(drainStart) + if secondDrain < time.Second { + secondDrain = time.Until(shutdownDeadline) + } + if secondDrain < time.Second { + secondDrain = time.Second + } + drain2, drain2Cancel := context.WithTimeout(context.Background(), secondDrain) + if drainErr := db.Queue.Drain(drain2); drainErr != nil && !errors.Is(drainErr, context.DeadlineExceeded) { + slog.Error("write queue drain", "err", drainErr) + } + drain2Cancel() + + waitRemain := time.Until(shutdownDeadline) + if waitRemain < 2*time.Second { + waitRemain = 2 * time.Second + } + waitCtx, waitCancel := context.WithTimeout(context.Background(), waitRemain) + _ = lnSrv.Wait(waitCtx) + waitCancel() return nil } func messageLoops(ctx context.Context, msgApp *message.App, db *store.DB, hashPool auth.HashPool, met *metrics.Registry) { msgApp.StartLoops(ctx) + defer msgApp.WaitLoops() t := time.NewTicker(15 * time.Second) defer t.Stop() sample := func() { diff --git a/cmd/nixmsg/serve_shutdown_test.go b/cmd/nixmsg/serve_shutdown_test.go new file mode 100644 index 0000000..6e69668 --- /dev/null +++ b/cmd/nixmsg/serve_shutdown_test.go @@ -0,0 +1,73 @@ +package main + +import ( + "context" + "sync" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/config" + "github.com/mochi-mqtt/server/v2/packets" +) + +func TestServeShutdownWithLiveMQTT(t *testing.T) { + dataDir := t.TempDir() + cfgPath := writeTestConfig(t, dataDir) + initAdminForTest(t, dataDir) + enableRegistration(t, dataDir, "uplink-code") + + cfg, err := config.Load(cfgPath) + if err != nil { + t.Fatal(err) + } + if vErr := cfg.Validate(); vErr != nil { + t.Fatal(vErr) + } + + ctx, cancel := context.WithCancel(context.Background()) + errCh := make(chan error, 1) + go func() { errCh <- runServe(ctx, cfg) }() + + addr := waitListenAddr(t, dataDir, 15*time.Second) + base := "http://" + addr + registerEP(t, base, "alice", "password12", "Alice") + registerEP(t, base, "bob", "password12", "Bob") + + ws := mqttSessionLogin(t, base, "alice", "password12") + tcp := mqttSessionLoginTCP(t, addr, "bob", "password12") + + cancel() + + var gotWG sync.WaitGroup + var wsReason, tcpReason byte + var wsOK, tcpOK bool + gotWG.Add(2) + go func() { + defer gotWG.Done() + wsReason, wsOK = ws.disconnectReason(12 * time.Second) + }() + go func() { + defer gotWG.Done() + tcpReason, tcpOK = tcp.disconnectReason(12 * time.Second) + }() + + select { + case err := <-errCh: + if err != nil { + t.Fatalf("runServe: %v", err) + } + case <-time.After(15 * time.Second): + t.Fatal("runServe did not return within 15s") + } + gotWG.Wait() + + want := packets.ErrServerShuttingDown.Code + if !wsOK || wsReason != want { + t.Fatalf("ws disconnect ok=%v reason=%#x want %#x", wsOK, wsReason, want) + } + if !tcpOK || tcpReason != want { + t.Fatalf("tcp disconnect ok=%v reason=%#x want %#x", tcpOK, tcpReason, want) + } + ws.Close() + tcp.Close() +} diff --git a/cmd/nixmsg/uplink_integration_test.go b/cmd/nixmsg/uplink_integration_test.go index bed1f90..44a7ed8 100644 --- a/cmd/nixmsg/uplink_integration_test.go +++ b/cmd/nixmsg/uplink_integration_test.go @@ -186,6 +186,7 @@ type mqttSess struct { inbox []map[string]any closed bool done chan struct{} + shutdownCh chan byte } type appResp struct { @@ -201,7 +202,19 @@ func mqttSessionLogin(t *testing.T, httpBase, endpointID, password string) *mqtt if err != nil { t.Fatalf("dial: %v", err) } - s := &mqttSess{t: t, mc: mc, endpointID: endpointID, pktID: 10, done: make(chan struct{})} + s := &mqttSess{t: t, mc: mc, endpointID: endpointID, pktID: 10, done: make(chan struct{}), shutdownCh: make(chan byte, 1)} + s.connectSubscribeHello(password) + go s.readLoop() + return s +} + +func mqttSessionLoginTCP(t *testing.T, addr, endpointID, password string) *mqttSess { + t.Helper() + mc, err := harness.DialMQTTTCP(addr, 10*time.Second) + if err != nil { + t.Fatalf("dial tcp: %v", err) + } + s := &mqttSess{t: t, mc: mc, endpointID: endpointID, pktID: 10, done: make(chan struct{}), shutdownCh: make(chan byte, 1)} s.connectSubscribeHello(password) go s.readLoop() return s @@ -344,6 +357,13 @@ func (s *mqttSess) handlePacket(raw []byte) map[string]any { switch typ { case packets.Puback, packets.Pingresp, packets.Suback: return nil + case packets.Disconnect: + reason := byte(0) + if _, n, err := decodeRemainingLength(raw[1:]); err == nil && 1+n < len(raw) { + reason = raw[1+n] + } + s.noteShutdown(reason) + return nil case packets.Publish: payload, err := decodePublishPayload(raw) if err != nil { @@ -446,6 +466,28 @@ func (s *mqttSess) takeMatching(pred func(map[string]any) bool) map[string]any { return nil } +func (s *mqttSess) noteShutdown(reason byte) { + if s.shutdownCh == nil { + return + } + select { + case s.shutdownCh <- reason: + default: + } +} + +func (s *mqttSess) disconnectReason(timeout time.Duration) (byte, bool) { + if s.shutdownCh == nil { + return 0, false + } + select { + case r := <-s.shutdownCh: + return r, true + case <-time.After(timeout): + return 0, false + } +} + func drainEvents(t *testing.T, s *mqttSess, d time.Duration) { t.Helper() deadline := time.Now().Add(d) diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index 7a197cf..d24cd62 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -12,6 +12,7 @@ services: # 本地/并行测试可加 container_name;正式部署可去掉 container_name: q4-nixmsg restart: unless-stopped + stop_grace_period: 30s command: ["serve"] environment: NIXMSG_CONFIG: /etc/nixmsg/config.yaml diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 211c06a..6c4ccc8 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -471,6 +471,15 @@ - 备选方案:metrics 失败锁定管理员 IP。 - 影响:锁定按真实客户端 IP;TLS 重连可 DidResume。 +### 复审修复 L-03 + +1. **有 MQTT 连接时停机** + - 原条款:DEVELOPMENT 7.8;issue #22。 + - 实际做法:listener 拆成 `StopAccept()`(关 TCP 监听并对 HTTP 调 Shutdown)和带超时的 `Wait(ctx)`(超时强关仍阻塞在 `OnMQTT` 的连接)。`serve` 在 `<-ctx.Done()` 后立刻调用 `stop()`;顺序为 StopAccept → 取消并等待消息循环 → Drain(10 秒)→ `brk.Shutdown`(5 秒)→ 用剩余时间再 Drain → `Wait` → `db.Close()`。compose `stop_grace_period: 30s`;OPS 写明 systemd `TimeoutStopSec` 至少 30 秒。未改 `PublishDown` 签名。 + - 原因:原先 `Close` 的 `wg.Wait` 会卡在裸 TCP/WS 的 `AttachTCP`/`AttachWS`,连 Drain 都走不到;`signal.NotifyContext` 的 stop 要等 `cmdServe` 返回才调用,卡住期间第二次 SIGTERM 被吞掉。 + - 备选方案:先强关全部 MQTT 再 HTTP Shutdown(会丢在途 HTTP)。 + - 影响:保持已登录 TCP/WS 客户端时 `runServe` 应在约 15 秒内返回,客户端收到 DISCONNECT `0x8B`。 + ## 消息 M ### M1 2026-09-30 diff --git a/docs/OPS.md b/docs/OPS.md index 4811923..9ffd592 100644 --- a/docs/OPS.md +++ b/docs/OPS.md @@ -122,6 +122,7 @@ curl -sS -H "Authorization: Bearer $NIXMSG_METRICS_TOKEN" http://127.0.0.1:7443/ - 容器以 uid `65532` 运行:挂载数据目录须可写;命名卷首次可 `docker run --rm -v <卷名>:/data busybox chown -R 65532:65532 /data`。 - 首次:`docker compose run --rm nixmsg admin init`,再 `up -d`。 +- Compose 示例已设 `stop_grace_period: 30s`。systemd 单元请设 `TimeoutStopSec=30`(或更长),以便进程先停接受、排空写队列并下发 MQTT DISCONNECT `0x8B`。第二次 Ctrl+C / SIGTERM 会按默认行为结束进程。 - 构建/推送 Task 目标见根目录 README(`q:docker-build` / `q:docker-push` / `q:docker-buildx`)。正式仓库推送在阶段 3。 ## 9. 验收与仍跳过的长时项 diff --git a/internal/listener/server.go b/internal/listener/server.go index c69dd72..a4eefd4 100644 --- a/internal/listener/server.go +++ b/internal/listener/server.go @@ -77,6 +77,11 @@ type Server struct { closed chan struct{} closeOnce sync.Once hsSem chan struct{} + + mqttMu sync.Mutex + mqttConns map[net.Conn]struct{} + waitOnce sync.Once + waitDone chan struct{} } // New 校验选项并准备证书;不开始监听。 @@ -100,11 +105,13 @@ func New(opts Options) (*Server, error) { hsLimit = defaultPreHandshakeLimit } s := &Server{ - opts: opts, - log: log, - proxies: ps, - closed: make(chan struct{}), - hsSem: make(chan struct{}, hsLimit), + opts: opts, + log: log, + proxies: ps, + closed: make(chan struct{}), + waitDone: make(chan struct{}), + mqttConns: make(map[net.Conn]struct{}), + hsSem: make(chan struct{}, hsLimit), } hasCert := opts.CertFile != "" && opts.KeyFile != "" if hasCert { @@ -183,7 +190,7 @@ func (s *Server) Start(ctx context.Context) error { go func() { <-ctx.Done() - _ = s.Close() + _ = s.StopAccept() }() return nil } @@ -205,8 +212,8 @@ func (s *Server) TLSConfig() *tls.Config { return s.certs.TLSConfig() } -// Close 停止接受并关闭 HTTP。 -func (s *Server) Close() error { +// StopAccept 停止接受新连接并对 HTTP 调用 Shutdown。不关闭已交给 OnMQTT 的连接,也不等待它们结束。 +func (s *Server) StopAccept() error { var first error s.closeOnce.Do(func() { close(s.closed) @@ -238,10 +245,74 @@ func (s *Server) Close() error { s.certs.Close() } }) - s.wg.Wait() return first } +// Wait 等待握手/HTTP/OnMQTT goroutine 结束。ctx 超时则强制关闭仍阻塞在 OnMQTT 的连接。 +func (s *Server) Wait(ctx context.Context) error { + s.waitOnce.Do(func() { + go func() { + s.wg.Wait() + close(s.waitDone) + }() + }) + if ctx == nil { + ctx = context.Background() + } + select { + case <-s.waitDone: + return nil + case <-ctx.Done(): + s.forceCloseMQTT() + select { + case <-s.waitDone: + case <-time.After(2 * time.Second): + } + return ctx.Err() + } +} + +// Close 停止接受,并限时等待剩余连接(超时则强关 OnMQTT 连接)。 +func (s *Server) Close() error { + first := s.StopAccept() + waitCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if err := s.Wait(waitCtx); err != nil && first == nil && !errors.Is(err, context.DeadlineExceeded) { + first = err + } + return first +} + +func (s *Server) trackMQTT(c net.Conn) { + if c == nil { + return + } + s.mqttMu.Lock() + s.mqttConns[c] = struct{}{} + s.mqttMu.Unlock() +} + +func (s *Server) untrackMQTT(c net.Conn) { + if c == nil { + return + } + s.mqttMu.Lock() + delete(s.mqttConns, c) + s.mqttMu.Unlock() +} + +func (s *Server) forceCloseMQTT() { + s.mqttMu.Lock() + conns := make([]net.Conn, 0, len(s.mqttConns)) + for c := range s.mqttConns { + conns = append(conns, c) + } + s.mqttMu.Unlock() + for _, c := range conns { + _ = c.Close() + } +} + func (s *Server) acceptLoop(ln net.Listener, isAdmin bool) { defer s.wg.Done() var delay time.Duration @@ -361,6 +432,8 @@ func (s *Server) handleConn(conn net.Conn, isAdmin bool) { httpLn.Enqueue(out) case KindMQTT: if s.opts.OnMQTT != nil { + s.trackMQTT(out) + defer s.untrackMQTT(out) _ = out.SetReadDeadline(time.Now().Add(s.handshakeTimeout())) s.opts.OnMQTT(out) } else { diff --git a/internal/listener/shutdown_test.go b/internal/listener/shutdown_test.go new file mode 100644 index 0000000..f40fef6 --- /dev/null +++ b/internal/listener/shutdown_test.go @@ -0,0 +1,81 @@ +package listener + +import ( + "context" + "net" + "net/http" + "testing" + "time" +) + +func TestStopAcceptDoesNotWaitForMQTT(t *testing.T) { + dir := t.TempDir() + held := make(chan net.Conn, 1) + mux := NewMux(RoleShared, Handlers{ + Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }), + }) + s, err := New(Options{ + Listen: "127.0.0.1:0", + DataDir: dir, + ClientHandler: mux, + AllowPlaintext: true, + OnMQTT: func(c net.Conn) { + held <- c + buf := make([]byte, 1) + for { + if _, err := c.Read(buf); err != nil { + break + } + } + _ = c.Close() + }, + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + if startErr := s.Start(ctx); startErr != nil { + t.Fatal(startErr) + } + + c, err := net.DialTimeout("tcp", s.ListenAddr(), 2*time.Second) + if err != nil { + t.Fatal(err) + } + defer func() { _ = c.Close() }() + if _, err := c.Write([]byte{0x10, 0x00}); err != nil { + t.Fatal(err) + } + select { + case <-held: + case <-time.After(3 * time.Second): + t.Fatal("OnMQTT not called") + } + + done := make(chan error, 1) + go func() { done <- s.StopAccept() }() + select { + case err := <-done: + if err != nil { + t.Fatalf("StopAccept: %v", err) + } + case <-time.After(3 * time.Second): + t.Fatal("StopAccept blocked on MQTT") + } + + waitCtx, waitCancel := context.WithTimeout(context.Background(), 300*time.Millisecond) + waitErr := s.Wait(waitCtx) + waitCancel() + if waitErr == nil { + t.Fatal("Wait returned before MQTT finished") + } + + waitCtx2, waitCancel2 := context.WithTimeout(context.Background(), 3*time.Second) + if err := s.Wait(waitCtx2); err != nil { + t.Fatalf("Wait after force-close: %v", err) + } + waitCancel2() +}