[优化] 增强 WebSocket 处理逻辑,添加上下文取消支持和关闭机制

[优化] 重构 WebSocket 处理逻辑,添加消息处理接口和心跳机制
This commit is contained in:
ryan
2026-06-02 19:25:56 +08:00
parent 4e58bdd85b
commit 4566fc1f53
9 changed files with 235 additions and 111 deletions
+55 -32
View File
@@ -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 {
@@ -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
+25 -56
View File
@@ -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) {
@@ -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()
}
+2
View File
@@ -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)
+47 -1
View File
@@ -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 {
+59
View File
@@ -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
+22 -22
View File
@@ -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) {
+11
View File
@@ -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()
}