fix: 停机先停接受并等待循环再断开 MQTT
This commit is contained in:
+55
-10
@@ -43,7 +43,20 @@ func cmdServe(_ []string) error {
|
|||||||
setupJSONLogger(cfg.Log)
|
setupJSONLogger(cfg.Log)
|
||||||
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
||||||
defer stop()
|
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 {
|
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()
|
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()
|
<-ctx.Done()
|
||||||
|
invokeServeStop(ctx)
|
||||||
|
shutdownDeadline := time.Now().Add(30 * time.Second)
|
||||||
|
_ = lnSrv.StopAccept()
|
||||||
loopCancel()
|
loopCancel()
|
||||||
// B-08:先对 MQTT 连接发 0x8B。HTTP Shutdown 与监听器完整停机顺序见 L-03。
|
<-loopsDone
|
||||||
shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
||||||
_ = brk.Shutdown(shutCtx)
|
drainBudget := 10 * time.Second
|
||||||
shutCancel()
|
drainStart := time.Now()
|
||||||
_ = lnSrv.Close()
|
drainCtx, drainCancel := context.WithTimeout(context.Background(), drainBudget)
|
||||||
drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
||||||
defer drainCancel()
|
|
||||||
if drainErr := db.Queue.Drain(drainCtx); drainErr != nil && !errors.Is(drainErr, context.DeadlineExceeded) {
|
if drainErr := db.Queue.Drain(drainCtx); drainErr != nil && !errors.Is(drainErr, context.DeadlineExceeded) {
|
||||||
slog.Error("write queue drain", "err", drainErr)
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func messageLoops(ctx context.Context, msgApp *message.App, db *store.DB, hashPool auth.HashPool, met *metrics.Registry) {
|
func messageLoops(ctx context.Context, msgApp *message.App, db *store.DB, hashPool auth.HashPool, met *metrics.Registry) {
|
||||||
msgApp.StartLoops(ctx)
|
msgApp.StartLoops(ctx)
|
||||||
|
defer msgApp.WaitLoops()
|
||||||
t := time.NewTicker(15 * time.Second)
|
t := time.NewTicker(15 * time.Second)
|
||||||
defer t.Stop()
|
defer t.Stop()
|
||||||
sample := func() {
|
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
|
inbox []map[string]any
|
||||||
closed bool
|
closed bool
|
||||||
done chan struct{}
|
done chan struct{}
|
||||||
|
shutdownCh chan byte
|
||||||
}
|
}
|
||||||
|
|
||||||
type appResp struct {
|
type appResp struct {
|
||||||
@@ -201,7 +202,19 @@ func mqttSessionLogin(t *testing.T, httpBase, endpointID, password string) *mqtt
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("dial: %v", err)
|
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)
|
s.connectSubscribeHello(password)
|
||||||
go s.readLoop()
|
go s.readLoop()
|
||||||
return s
|
return s
|
||||||
@@ -344,6 +357,13 @@ func (s *mqttSess) handlePacket(raw []byte) map[string]any {
|
|||||||
switch typ {
|
switch typ {
|
||||||
case packets.Puback, packets.Pingresp, packets.Suback:
|
case packets.Puback, packets.Pingresp, packets.Suback:
|
||||||
return nil
|
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:
|
case packets.Publish:
|
||||||
payload, err := decodePublishPayload(raw)
|
payload, err := decodePublishPayload(raw)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -446,6 +466,28 @@ func (s *mqttSess) takeMatching(pred func(map[string]any) bool) map[string]any {
|
|||||||
return nil
|
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) {
|
func drainEvents(t *testing.T, s *mqttSess, d time.Duration) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
deadline := time.Now().Add(d)
|
deadline := time.Now().Add(d)
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ services:
|
|||||||
# 本地/并行测试可加 container_name;正式部署可去掉
|
# 本地/并行测试可加 container_name;正式部署可去掉
|
||||||
container_name: q4-nixmsg
|
container_name: q4-nixmsg
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
|
stop_grace_period: 30s
|
||||||
command: ["serve"]
|
command: ["serve"]
|
||||||
environment:
|
environment:
|
||||||
NIXMSG_CONFIG: /etc/nixmsg/config.yaml
|
NIXMSG_CONFIG: /etc/nixmsg/config.yaml
|
||||||
|
|||||||
@@ -471,6 +471,15 @@
|
|||||||
- 备选方案:metrics 失败锁定管理员 IP。
|
- 备选方案:metrics 失败锁定管理员 IP。
|
||||||
- 影响:锁定按真实客户端 IP;TLS 重连可 DidResume。
|
- 影响:锁定按真实客户端 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
|
## 消息 M
|
||||||
|
|
||||||
### M1 2026-09-30
|
### M1 2026-09-30
|
||||||
|
|||||||
@@ -122,6 +122,7 @@ curl -sS -H "Authorization: Bearer $NIXMSG_METRICS_TOKEN" http://127.0.0.1:7443/
|
|||||||
- 容器以 uid `65532` 运行:挂载数据目录须可写;命名卷首次可
|
- 容器以 uid `65532` 运行:挂载数据目录须可写;命名卷首次可
|
||||||
`docker run --rm -v <卷名>:/data busybox chown -R 65532:65532 /data`。
|
`docker run --rm -v <卷名>:/data busybox chown -R 65532:65532 /data`。
|
||||||
- 首次:`docker compose run --rm nixmsg admin init`,再 `up -d`。
|
- 首次:`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。
|
- 构建/推送 Task 目标见根目录 README(`q:docker-build` / `q:docker-push` / `q:docker-buildx`)。正式仓库推送在阶段 3。
|
||||||
|
|
||||||
## 9. 验收与仍跳过的长时项
|
## 9. 验收与仍跳过的长时项
|
||||||
|
|||||||
@@ -77,6 +77,11 @@ type Server struct {
|
|||||||
closed chan struct{}
|
closed chan struct{}
|
||||||
closeOnce sync.Once
|
closeOnce sync.Once
|
||||||
hsSem chan struct{}
|
hsSem chan struct{}
|
||||||
|
|
||||||
|
mqttMu sync.Mutex
|
||||||
|
mqttConns map[net.Conn]struct{}
|
||||||
|
waitOnce sync.Once
|
||||||
|
waitDone chan struct{}
|
||||||
}
|
}
|
||||||
|
|
||||||
// New 校验选项并准备证书;不开始监听。
|
// New 校验选项并准备证书;不开始监听。
|
||||||
@@ -100,11 +105,13 @@ func New(opts Options) (*Server, error) {
|
|||||||
hsLimit = defaultPreHandshakeLimit
|
hsLimit = defaultPreHandshakeLimit
|
||||||
}
|
}
|
||||||
s := &Server{
|
s := &Server{
|
||||||
opts: opts,
|
opts: opts,
|
||||||
log: log,
|
log: log,
|
||||||
proxies: ps,
|
proxies: ps,
|
||||||
closed: make(chan struct{}),
|
closed: make(chan struct{}),
|
||||||
hsSem: make(chan struct{}, hsLimit),
|
waitDone: make(chan struct{}),
|
||||||
|
mqttConns: make(map[net.Conn]struct{}),
|
||||||
|
hsSem: make(chan struct{}, hsLimit),
|
||||||
}
|
}
|
||||||
hasCert := opts.CertFile != "" && opts.KeyFile != ""
|
hasCert := opts.CertFile != "" && opts.KeyFile != ""
|
||||||
if hasCert {
|
if hasCert {
|
||||||
@@ -183,7 +190,7 @@ func (s *Server) Start(ctx context.Context) error {
|
|||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
<-ctx.Done()
|
<-ctx.Done()
|
||||||
_ = s.Close()
|
_ = s.StopAccept()
|
||||||
}()
|
}()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -205,8 +212,8 @@ func (s *Server) TLSConfig() *tls.Config {
|
|||||||
return s.certs.TLSConfig()
|
return s.certs.TLSConfig()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close 停止接受并关闭 HTTP。
|
// StopAccept 停止接受新连接并对 HTTP 调用 Shutdown。不关闭已交给 OnMQTT 的连接,也不等待它们结束。
|
||||||
func (s *Server) Close() error {
|
func (s *Server) StopAccept() error {
|
||||||
var first error
|
var first error
|
||||||
s.closeOnce.Do(func() {
|
s.closeOnce.Do(func() {
|
||||||
close(s.closed)
|
close(s.closed)
|
||||||
@@ -238,10 +245,74 @@ func (s *Server) Close() error {
|
|||||||
s.certs.Close()
|
s.certs.Close()
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
s.wg.Wait()
|
|
||||||
return first
|
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) {
|
func (s *Server) acceptLoop(ln net.Listener, isAdmin bool) {
|
||||||
defer s.wg.Done()
|
defer s.wg.Done()
|
||||||
var delay time.Duration
|
var delay time.Duration
|
||||||
@@ -361,6 +432,8 @@ func (s *Server) handleConn(conn net.Conn, isAdmin bool) {
|
|||||||
httpLn.Enqueue(out)
|
httpLn.Enqueue(out)
|
||||||
case KindMQTT:
|
case KindMQTT:
|
||||||
if s.opts.OnMQTT != nil {
|
if s.opts.OnMQTT != nil {
|
||||||
|
s.trackMQTT(out)
|
||||||
|
defer s.untrackMQTT(out)
|
||||||
_ = out.SetReadDeadline(time.Now().Add(s.handshakeTimeout()))
|
_ = out.SetReadDeadline(time.Now().Add(s.handshakeTimeout()))
|
||||||
s.opts.OnMQTT(out)
|
s.opts.OnMQTT(out)
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user