fix: 停机先停接受并等待循环再断开 MQTT

This commit is contained in:
Nixevol
2026-09-30 16:22:44 +08:00
parent 6a65ab8593
commit b40ef5c548
8 changed files with 345 additions and 20 deletions
+82 -9
View File
@@ -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 {
+81
View File
@@ -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()
}