From 4566fc1f53aad14a53da1d0624c74eced9e22975 Mon Sep 17 00:00:00 2001 From: ryan Date: Tue, 2 Jun 2026 19:25:56 +0800 Subject: [PATCH] =?UTF-8?q?[=E4=BC=98=E5=8C=96]=20=E5=A2=9E=E5=BC=BA=20Web?= =?UTF-8?q?Socket=20=E5=A4=84=E7=90=86=E9=80=BB=E8=BE=91=EF=BC=8C=E6=B7=BB?= =?UTF-8?q?=E5=8A=A0=E4=B8=8A=E4=B8=8B=E6=96=87=E5=8F=96=E6=B6=88=E6=94=AF?= =?UTF-8?q?=E6=8C=81=E5=92=8C=E5=85=B3=E9=97=AD=E6=9C=BA=E5=88=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit [优化] 重构 WebSocket 处理逻辑,添加消息处理接口和心跳机制 --- openflare_agent/internal/agent/runner.go | 87 +++++++++++++-------- openflare_agent/internal/wsclient/client.go | 7 ++ openflare_relay/internal/relay/runner.go | 81 ++++++------------- openflare_relay/internal/wsclient/client.go | 7 ++ openflare_server/main.go | 2 + openflare_server/service/ws_hub.go | 48 +++++++++++- openflare_server/utils/wsclient/client.go | 59 ++++++++++++++ openflared/internal/flared/runner.go | 44 +++++------ openflared/internal/wsclient/client.go | 11 +++ 9 files changed, 235 insertions(+), 111 deletions(-) diff --git a/openflare_agent/internal/agent/runner.go b/openflare_agent/internal/agent/runner.go index 598c93d3..54ccee78 100644 --- a/openflare_agent/internal/agent/runner.go +++ b/openflare_agent/internal/agent/runner.go @@ -12,6 +12,7 @@ import ( "openflare-agent/internal/observability" "openflare-agent/internal/protocol" "openflare-agent/internal/state" + "openflare-agent/internal/wsclient" ) type HeartbeatService interface { @@ -211,53 +212,75 @@ func (r *Runner) startWebSocket(ctx context.Context, nodeID string) (<-chan erro return done, nil } +type agentWSHandler struct { + runner *Runner + conn protocol.WebSocketConnection + nodeID string + statusTicker *time.Ticker +} + +func (h *agentWSHandler) OnConnect(ctx context.Context) error { + return h.runner.sendWebSocketStatus(ctx, h.nodeID, h.conn) +} + +func (h *agentWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error { + var payloadBytes []byte + if msg.Payload != nil { + payloadBytes = []byte(msg.Payload) + } + protoMsg := protocol.WSMessage{ + Type: msg.Type, + Payload: payloadBytes, + } + changed, err := h.runner.handleWebSocketMessage(ctx, protoMsg, h.conn) + if err != nil { + return err + } + if changed { + h.statusTicker.Reset(h.runner.Config.HeartbeatInterval.Duration()) + } + return nil +} + +func (h *agentWSHandler) OnClose(err error) { + slog.Error("agent ws receive failed", "error", err) +} + func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error { slog.Debug("agent ws connected", "url", conn.URL(), "node_id", nodeID) statusTicker := time.NewTicker(r.Config.HeartbeatInterval.Duration()) defer statusTicker.Stop() - messages := make(chan protocol.WSMessage, 8) - readDone := make(chan error, 1) + childCtx, cancel := context.WithCancel(ctx) + defer cancel() + + // Start status ticker sender in background go func() { for { - message, err := conn.Receive() - if err != nil { - readDone <- err - return - } select { - case messages <- message: - case <-ctx.Done(): - readDone <- ctx.Err() + case <-childCtx.Done(): return + case <-statusTicker.C: + if err := r.sendWebSocketStatus(childCtx, nodeID, conn); err != nil { + slog.Error("agent ws send status failed", "error", err) + _ = conn.Close() + return + } } } }() - if err := r.sendWebSocketStatus(ctx, nodeID, conn); err != nil { - return err + wsConn, ok := conn.(*wsclient.Connection) + if !ok { + return errors.New("invalid websocket connection type") } - for { - select { - case <-ctx.Done(): - return ctx.Err() - case err := <-readDone: - return err - case <-statusTicker.C: - if err := r.sendWebSocketStatus(ctx, nodeID, conn); err != nil { - return err - } - case message := <-messages: - changed, err := r.handleWebSocketMessage(ctx, message, conn) - if err != nil { - return err - } - if changed { - statusTicker.Reset(r.Config.HeartbeatInterval.Duration()) - } - } - } + return wsConn.RunReceiveLoop(childCtx, &agentWSHandler{ + runner: r, + conn: conn, + nodeID: nodeID, + statusTicker: statusTicker, + }) } func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error { diff --git a/openflare_agent/internal/wsclient/client.go b/openflare_agent/internal/wsclient/client.go index 808294b9..2e85a96e 100644 --- a/openflare_agent/internal/wsclient/client.go +++ b/openflare_agent/internal/wsclient/client.go @@ -8,6 +8,9 @@ import ( shared "openflare/utils/wsclient" ) +type WSMessage = shared.WSMessage +type MessageHandler = shared.MessageHandler + type Client struct { sharedClient *shared.Client } @@ -67,6 +70,10 @@ func (conn *Connection) Receive() (protocol.WSMessage, error) { return message, nil } +func (conn *Connection) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error { + return conn.sharedConn.RunReceiveLoop(ctx, handler) +} + func (conn *Connection) Close() error { if conn == nil || conn.sharedConn == nil { return nil diff --git a/openflare_relay/internal/relay/runner.go b/openflare_relay/internal/relay/runner.go index 2e1df355..752c673c 100644 --- a/openflare_relay/internal/relay/runner.go +++ b/openflare_relay/internal/relay/runner.go @@ -51,66 +51,35 @@ func (r *Runner) Run(ctx context.Context) error { } } -func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) { - // Send pings at 2× heartbeat interval to keep the server-side read deadline - // from expiring (server closes the WS if no data arrives within ~30 s). - pingInterval := r.Config.HeartbeatInterval.Duration() * 2 - pingTicker := time.NewTicker(pingInterval) - defer pingTicker.Stop() +type relayWSHandler struct { + runner *Runner +} - messages := make(chan service.WSMessage, 8) - readDone := make(chan error, 1) - go func() { - for { - msg, err := conn.Receive() - if err != nil { - readDone <- err - return - } - select { - case messages <- msg: - case <-ctx.Done(): - readDone <- ctx.Err() - return - } - } - }() +func (h *relayWSHandler) OnConnect(ctx context.Context) error { + return nil +} - for { - select { - case <-ctx.Done(): - return - case err := <-readDone: - slog.Error("relay ws receive failed", "error", err) - return - case <-pingTicker.C: - if err := conn.SendPing(); err != nil { - slog.Error("relay ws send ping failed", "error", err) - return - } - case msg := <-messages: - switch msg.Type { - case "ping": - _ = conn.SendPong() - case "pong": - slog.Debug("relay ws pong received") - case "relay_config": - payloadBytes, ok := msg.Payload.(json.RawMessage) - if !ok { - slog.Error("invalid relay_config payload type") - continue - } - var cfg service.RelayConfig - if err := json.Unmarshal(payloadBytes, &cfg); err != nil { - slog.Error("failed to unmarshal relay_config", "error", err) - continue - } - r.FrpsManager.UpdateConfig(&cfg) - default: - slog.Debug("ignored unknown ws message type", "type", msg.Type) - } +func (h *relayWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error { + switch msg.Type { + case "relay_config": + var cfg service.RelayConfig + if err := json.Unmarshal(msg.Payload, &cfg); err != nil { + slog.Error("failed to unmarshal relay_config", "error", err) + return nil } + h.runner.FrpsManager.UpdateConfig(&cfg) + default: + slog.Debug("ignored unknown ws message type", "type", msg.Type) } + return nil +} + +func (h *relayWSHandler) OnClose(err error) { + slog.Error("relay ws receive failed", "error", err) +} + +func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) { + _ = conn.RunReceiveLoop(ctx, &relayWSHandler{runner: r}) } func (r *Runner) sleepContext(ctx context.Context, d time.Duration) { diff --git a/openflare_relay/internal/wsclient/client.go b/openflare_relay/internal/wsclient/client.go index 434ee580..e2438ea6 100644 --- a/openflare_relay/internal/wsclient/client.go +++ b/openflare_relay/internal/wsclient/client.go @@ -9,6 +9,9 @@ import ( shared "openflare/utils/wsclient" ) +type WSMessage = shared.WSMessage +type MessageHandler = shared.MessageHandler + type Client struct { sharedClient *shared.Client } @@ -63,6 +66,10 @@ func (conn *Connection) Receive() (service.WSMessage, error) { }, nil } +func (conn *Connection) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error { + return conn.sharedConn.RunReceiveLoop(ctx, handler) +} + func (conn *Connection) Close() error { return conn.sharedConn.Close() } diff --git a/openflare_server/main.go b/openflare_server/main.go index 91c6e946..7bff9b25 100644 --- a/openflare_server/main.go +++ b/openflare_server/main.go @@ -49,6 +49,8 @@ func main() { gin.SetMode(gin.ReleaseMode) } // Initialize SQL Database + defer service.ShutdownWSHubs() + err := model.InitDB() if err != nil { slog.Error("initialize database failed", "error", err) diff --git a/openflare_server/service/ws_hub.go b/openflare_server/service/ws_hub.go index 8da8ba18..7de529f6 100644 --- a/openflare_server/service/ws_hub.go +++ b/openflare_server/service/ws_hub.go @@ -3,6 +3,7 @@ package service import ( "log/slog" "sync" + "time" ) type WSMessage struct { @@ -65,13 +66,58 @@ type WSHub struct { name string mu sync.RWMutex clients map[string]*WSClient + done chan struct{} } func NewWSHub(name string) *WSHub { - return &WSHub{ + h := &WSHub{ name: name, clients: make(map[string]*WSClient), + done: make(chan struct{}), } + go h.startPingLoop() + return h +} + +func (h *WSHub) Close() { + close(h.done) +} + +func (h *WSHub) startPingLoop() { + ticker := time.NewTicker(10 * time.Second) + defer ticker.Stop() + for { + select { + case <-h.done: + return + case <-ticker.C: + h.mu.RLock() + if len(h.clients) == 0 { + h.mu.RUnlock() + continue + } + clients := make([]*WSClient, 0, len(h.clients)) + for _, client := range h.clients { + clients = append(clients, client) + } + h.mu.RUnlock() + + for _, client := range clients { + if !client.Send(WSMessage{ + Type: "ping", + }) { + slog.Warn("ws client send ping failed, queue full, disconnecting", "hub", h.name, "id", client.id) + h.Disconnect(client.id) + } + } + } + } +} + +func ShutdownWSHubs() { + DefaultAgentWSHub.Close() + DefaultFlaredWSHub.Close() + DefaultRelayWSHub.Close() } func (h *WSHub) Register(id string) *WSClient { diff --git a/openflare_server/utils/wsclient/client.go b/openflare_server/utils/wsclient/client.go index fd4dc48f..1260f85d 100644 --- a/openflare_server/utils/wsclient/client.go +++ b/openflare_server/utils/wsclient/client.go @@ -2,6 +2,7 @@ package wsclient import ( "context" + "encoding/json" "errors" "log/slog" "net" @@ -25,6 +26,17 @@ type Client struct { cfg Config } +type WSMessage struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` +} + +type MessageHandler interface { + OnConnect(ctx context.Context) error + HandleMessage(ctx context.Context, msg WSMessage) error + OnClose(err error) +} + type Connection struct { Conn *websocket.Conn URL string @@ -155,6 +167,53 @@ func websocketReadTimeout(requestTimeout time.Duration) time.Duration { return timeout } +func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandler) error { + doneChan := make(chan struct{}) + defer close(doneChan) + + go func() { + select { + case <-ctx.Done(): + _ = conn.Close() + case <-doneChan: + } + }() + + if err := handler.OnConnect(ctx); err != nil { + handler.OnClose(err) + return err + } + + for { + select { + case <-ctx.Done(): + return ctx.Err() + default: + } + + var raw WSMessage + if err := conn.Receive(&raw); err != nil { + handler.OnClose(err) + return err + } + + switch raw.Type { + case "ping": + slog.Debug("ws received ping from server, replying with pong") + if err := conn.SendMessage("pong", nil); err != nil { + slog.Error("ws send pong response failed", "error", err) + } + case "pong": + slog.Debug("ws received pong response from server") + default: + if err := handler.HandleMessage(ctx, raw); err != nil { + slog.Error("ws handler failed to process message", "type", raw.Type, "error", err) + return err + } + } + } +} + func (conn *Connection) Close() error { if conn == nil || conn.Conn == nil { return nil diff --git a/openflared/internal/flared/runner.go b/openflared/internal/flared/runner.go index 0555f5db..ef2a958d 100644 --- a/openflared/internal/flared/runner.go +++ b/openflared/internal/flared/runner.go @@ -49,31 +49,31 @@ func (r *Runner) Run(ctx context.Context) error { } } -func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) { - for { - select { - case <-ctx.Done(): - return - default: - } +type flaredWSHandler struct { + runner *Runner +} - msg, err := conn.Receive() - if err != nil { - slog.Error("flared ws receive failed", "error", err) - return - } +func (h *flaredWSHandler) OnConnect(ctx context.Context) error { + return nil +} - switch msg.Type { - case "ping": - _ = conn.SendPong() - case "active_config": - // Server notifies there is a new config available - slog.Info("received config update notification from server") - r.SyncService.Trigger() - default: - slog.Debug("ignored unknown ws message type", "type", msg.Type) - } +func (h *flaredWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error { + switch msg.Type { + case "active_config": + slog.Info("received config update notification from server") + h.runner.SyncService.Trigger() + default: + slog.Debug("ignored unknown ws message type", "type", msg.Type) } + return nil +} + +func (h *flaredWSHandler) OnClose(err error) { + slog.Error("flared ws receive failed", "error", err) +} + +func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) { + _ = conn.RunReceiveLoop(ctx, &flaredWSHandler{runner: r}) } func (r *Runner) sleepContext(ctx context.Context, d time.Duration) { diff --git a/openflared/internal/wsclient/client.go b/openflared/internal/wsclient/client.go index e567621d..40564d59 100644 --- a/openflared/internal/wsclient/client.go +++ b/openflared/internal/wsclient/client.go @@ -9,6 +9,9 @@ import ( shared "openflare/utils/wsclient" ) +type WSMessage = shared.WSMessage +type MessageHandler = shared.MessageHandler + type Client struct { sharedClient *shared.Client } @@ -41,6 +44,10 @@ func (c *Client) Connect(ctx context.Context) (*Connection, error) { return &Connection{sharedConn: conn}, nil } +func (conn *Connection) SendPing() error { + return conn.sharedConn.SendMessage("ping", nil) +} + func (conn *Connection) SendPong() error { return conn.sharedConn.SendMessage("pong", nil) } @@ -59,6 +66,10 @@ func (conn *Connection) Receive() (service.WSMessage, error) { }, nil } +func (conn *Connection) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error { + return conn.sharedConn.RunReceiveLoop(ctx, handler) +} + func (conn *Connection) Close() error { return conn.sharedConn.Close() }