From 5d49d554567343ac5ea32faad18e175fa1000ce1 Mon Sep 17 00:00:00 2001 From: HuaiYJ Date: Mon, 24 Aug 2026 16:21:14 +0800 Subject: [PATCH] fix: close active data channel connections --- internal/datachan/server/server.go | 39 ++++++++++++++++++++++++++++++ 1 file changed, 39 insertions(+) diff --git a/internal/datachan/server/server.go b/internal/datachan/server/server.go index f37fe3d..701033f 100644 --- a/internal/datachan/server/server.go +++ b/internal/datachan/server/server.go @@ -40,6 +40,9 @@ type Server struct { listener net.Listener cancel context.CancelFunc wg sync.WaitGroup + mu sync.Mutex + conns map[net.Conn]struct{} + stopping bool log *zap.Logger reporter *data.Reporter } @@ -51,6 +54,7 @@ func NewServer(path string, log *zap.Logger, reporter *data.Reporter) *Server { Path: path, log: log, reporter: reporter, + conns: make(map[net.Conn]struct{}), } } @@ -68,6 +72,9 @@ func (s *Server) Start() error { return fmt.Errorf("failed to listen on %s: %w", s.Path, err) } s.listener = listener + s.mu.Lock() + s.stopping = false + s.mu.Unlock() ctx, cancel := context.WithCancel(context.Background()) s.cancel = cancel @@ -99,12 +106,23 @@ func removeStaleSocket(path string) error { // Stop 优雅关闭服务器: // 取消 acceptLoop 上下文、关闭监听、等待所有连接处理 goroutine 退出,并删除套接字文件。 func (s *Server) Stop() { + s.mu.Lock() + s.stopping = true + connections := make([]net.Conn, 0, len(s.conns)) + for conn := range s.conns { + connections = append(connections, conn) + } + s.mu.Unlock() + if s.cancel != nil { s.cancel() } if s.listener != nil { s.listener.Close() } + for _, conn := range connections { + conn.Close() + } s.wg.Wait() os.Remove(s.Path) s.log.Info("Unix socket server stopped") @@ -126,6 +144,10 @@ func (s *Server) acceptLoop(ctx context.Context) { continue } } + if !s.trackConn(conn) { + conn.Close() + continue + } s.wg.Add(1) go s.handleConn(ctx, conn) } @@ -136,6 +158,7 @@ func (s *Server) acceptLoop(ctx context.Context) { // 报文长度非法(0 或 > MaxMessageLen)或 ctx 取消时返回。 func (s *Server) handleConn(ctx context.Context, conn net.Conn) { defer s.wg.Done() + defer s.untrackConn(conn) defer conn.Close() for { @@ -190,3 +213,19 @@ func (s *Server) handleConn(ctx context.Context, conn net.Conn) { } } } + +func (s *Server) trackConn(conn net.Conn) bool { + s.mu.Lock() + defer s.mu.Unlock() + if s.stopping { + return false + } + s.conns[conn] = struct{}{} + return true +} + +func (s *Server) untrackConn(conn net.Conn) { + s.mu.Lock() + delete(s.conns, conn) + s.mu.Unlock() +} -- Gitee