diff --git a/go-backend/internal/ws/server.go b/go-backend/internal/ws/server.go index e1d35da..f27e7bc 100644 --- a/go-backend/internal/ws/server.go +++ b/go-backend/internal/ws/server.go @@ -54,6 +54,12 @@ type pendingRequest struct { ch chan CommandResult } +const ( + wsPingPeriod = 15 * time.Second + wsPongWait = 45 * time.Second + wsWriteWait = 5 * time.Second +) + type CommandResult struct { Type string `json:"type"` Success bool `json:"success"` @@ -120,12 +126,19 @@ func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) { return } cw := &connWrap{conn: conn} + _ = conn.SetReadDeadline(time.Now().Add(wsPongWait)) + conn.SetPongHandler(func(string) error { + return conn.SetReadDeadline(time.Now().Add(wsPongWait)) + }) + done := make(chan struct{}) + go startKeepalive(cw, done) s.mu.Lock() s.admins[cw] = struct{}{} s.mu.Unlock() defer func() { + close(done) s.mu.Lock() delete(s.admins, cw) s.mu.Unlock() @@ -145,6 +158,12 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64 return } cw := &connWrap{conn: conn} + _ = conn.SetReadDeadline(time.Now().Add(wsPongWait)) + conn.SetPongHandler(func(string) error { + return conn.SetReadDeadline(time.Now().Add(wsPongWait)) + }) + done := make(chan struct{}) + go startKeepalive(cw, done) version := r.URL.Query().Get("version") httpVal := parseIntDefault(r.URL.Query().Get("http"), 0) @@ -165,6 +184,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64 s.broadcastStatus(nodeID, 1) defer func() { + close(done) needOfflineBroadcast := false s.mu.Lock() current, ok := s.nodes[nodeID] @@ -442,3 +462,27 @@ func parseIntDefault(v string, fallback int) int { } return x } + +func startKeepalive(cw *connWrap, done <-chan struct{}) { + if cw == nil || cw.conn == nil { + return + } + ticker := time.NewTicker(wsPingPeriod) + defer ticker.Stop() + + for { + select { + case <-done: + return + case <-ticker.C: + cw.mu.Lock() + _ = cw.conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) + err := cw.conn.WriteMessage(websocket.PingMessage, nil) + cw.mu.Unlock() + if err != nil { + _ = cw.conn.Close() + return + } + } + } +} diff --git a/go-gost/x/socket/websocket_reporter.go b/go-gost/x/socket/websocket_reporter.go index 8376ce1..22bb362 100644 --- a/go-gost/x/socket/websocket_reporter.go +++ b/go-gost/x/socket/websocket_reporter.go @@ -91,6 +91,11 @@ type TcpPingResponse struct { RequestId string `json:"requestId,omitempty"` } +const ( + reporterReadWait = 60 * time.Second + reporterWriteWait = 5 * time.Second +) + type WebSocketReporter struct { url string addr string // 保存服务器地址 @@ -243,6 +248,14 @@ func (w *WebSocketReporter) connect() error { w.conn = conn w.connected = true + _ = conn.SetReadDeadline(time.Now().Add(reporterReadWait)) + conn.SetPingHandler(func(appData string) error { + _ = conn.SetReadDeadline(time.Now().Add(reporterReadWait)) + return conn.WriteControl(websocket.PongMessage, []byte(appData), time.Now().Add(reporterWriteWait)) + }) + conn.SetPongHandler(func(string) error { + return conn.SetReadDeadline(time.Now().Add(reporterReadWait)) + }) // 设置关闭处理器来检测连接状态 w.conn.SetCloseHandler(func(code int, text string) error { @@ -383,7 +396,7 @@ func (w *WebSocketReporter) receiveMessages() { } // 设置读取超时 - conn.SetReadDeadline(time.Now().Add(30 * time.Second)) + conn.SetReadDeadline(time.Now().Add(reporterReadWait)) messageType, message, err := conn.ReadMessage() if err != nil {