fix: 停机先停接受并等待循环再断开 MQTT
This commit is contained in:
+55
-10
@@ -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() {
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user