fix: 停机先停接受并等待循环再断开 MQTT
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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