mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 17:56:37 +08:00
[优化] 增强 WebSocket 处理逻辑,添加上下文取消支持和关闭机制
[优化] 重构 WebSocket 处理逻辑,添加消息处理接口和心跳机制
This commit is contained in:
@@ -12,6 +12,7 @@ import (
|
|||||||
"openflare-agent/internal/observability"
|
"openflare-agent/internal/observability"
|
||||||
"openflare-agent/internal/protocol"
|
"openflare-agent/internal/protocol"
|
||||||
"openflare-agent/internal/state"
|
"openflare-agent/internal/state"
|
||||||
|
"openflare-agent/internal/wsclient"
|
||||||
)
|
)
|
||||||
|
|
||||||
type HeartbeatService interface {
|
type HeartbeatService interface {
|
||||||
@@ -211,53 +212,75 @@ func (r *Runner) startWebSocket(ctx context.Context, nodeID string) (<-chan erro
|
|||||||
return done, nil
|
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 {
|
func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
|
||||||
slog.Debug("agent ws connected", "url", conn.URL(), "node_id", nodeID)
|
slog.Debug("agent ws connected", "url", conn.URL(), "node_id", nodeID)
|
||||||
statusTicker := time.NewTicker(r.Config.HeartbeatInterval.Duration())
|
statusTicker := time.NewTicker(r.Config.HeartbeatInterval.Duration())
|
||||||
defer statusTicker.Stop()
|
defer statusTicker.Stop()
|
||||||
|
|
||||||
messages := make(chan protocol.WSMessage, 8)
|
childCtx, cancel := context.WithCancel(ctx)
|
||||||
readDone := make(chan error, 1)
|
defer cancel()
|
||||||
|
|
||||||
|
// Start status ticker sender in background
|
||||||
go func() {
|
go func() {
|
||||||
for {
|
for {
|
||||||
message, err := conn.Receive()
|
|
||||||
if err != nil {
|
|
||||||
readDone <- err
|
|
||||||
return
|
|
||||||
}
|
|
||||||
select {
|
select {
|
||||||
case messages <- message:
|
case <-childCtx.Done():
|
||||||
case <-ctx.Done():
|
|
||||||
readDone <- ctx.Err()
|
|
||||||
return
|
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 {
|
wsConn, ok := conn.(*wsclient.Connection)
|
||||||
return err
|
if !ok {
|
||||||
|
return errors.New("invalid websocket connection type")
|
||||||
}
|
}
|
||||||
|
|
||||||
for {
|
return wsConn.RunReceiveLoop(childCtx, &agentWSHandler{
|
||||||
select {
|
runner: r,
|
||||||
case <-ctx.Done():
|
conn: conn,
|
||||||
return ctx.Err()
|
nodeID: nodeID,
|
||||||
case err := <-readDone:
|
statusTicker: statusTicker,
|
||||||
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())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
|
func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
|
||||||
|
|||||||
@@ -8,6 +8,9 @@ import (
|
|||||||
shared "openflare/utils/wsclient"
|
shared "openflare/utils/wsclient"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type WSMessage = shared.WSMessage
|
||||||
|
type MessageHandler = shared.MessageHandler
|
||||||
|
|
||||||
type Client struct {
|
type Client struct {
|
||||||
sharedClient *shared.Client
|
sharedClient *shared.Client
|
||||||
}
|
}
|
||||||
@@ -67,6 +70,10 @@ func (conn *Connection) Receive() (protocol.WSMessage, error) {
|
|||||||
return message, nil
|
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 {
|
func (conn *Connection) Close() error {
|
||||||
if conn == nil || conn.sharedConn == nil {
|
if conn == nil || conn.sharedConn == nil {
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -51,66 +51,35 @@ func (r *Runner) Run(ctx context.Context) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) {
|
type relayWSHandler struct {
|
||||||
// Send pings at 2× heartbeat interval to keep the server-side read deadline
|
runner *Runner
|
||||||
// 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()
|
|
||||||
|
|
||||||
messages := make(chan service.WSMessage, 8)
|
func (h *relayWSHandler) OnConnect(ctx context.Context) error {
|
||||||
readDone := make(chan error, 1)
|
return nil
|
||||||
go func() {
|
}
|
||||||
for {
|
|
||||||
msg, err := conn.Receive()
|
|
||||||
if err != nil {
|
|
||||||
readDone <- err
|
|
||||||
return
|
|
||||||
}
|
|
||||||
select {
|
|
||||||
case messages <- msg:
|
|
||||||
case <-ctx.Done():
|
|
||||||
readDone <- ctx.Err()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
for {
|
func (h *relayWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
|
||||||
select {
|
switch msg.Type {
|
||||||
case <-ctx.Done():
|
case "relay_config":
|
||||||
return
|
var cfg service.RelayConfig
|
||||||
case err := <-readDone:
|
if err := json.Unmarshal(msg.Payload, &cfg); err != nil {
|
||||||
slog.Error("relay ws receive failed", "error", err)
|
slog.Error("failed to unmarshal relay_config", "error", err)
|
||||||
return
|
return nil
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
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) {
|
func (r *Runner) sleepContext(ctx context.Context, d time.Duration) {
|
||||||
|
|||||||
@@ -9,6 +9,9 @@ import (
|
|||||||
shared "openflare/utils/wsclient"
|
shared "openflare/utils/wsclient"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type WSMessage = shared.WSMessage
|
||||||
|
type MessageHandler = shared.MessageHandler
|
||||||
|
|
||||||
type Client struct {
|
type Client struct {
|
||||||
sharedClient *shared.Client
|
sharedClient *shared.Client
|
||||||
}
|
}
|
||||||
@@ -63,6 +66,10 @@ func (conn *Connection) Receive() (service.WSMessage, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error {
|
||||||
|
return conn.sharedConn.RunReceiveLoop(ctx, handler)
|
||||||
|
}
|
||||||
|
|
||||||
func (conn *Connection) Close() error {
|
func (conn *Connection) Close() error {
|
||||||
return conn.sharedConn.Close()
|
return conn.sharedConn.Close()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -49,6 +49,8 @@ func main() {
|
|||||||
gin.SetMode(gin.ReleaseMode)
|
gin.SetMode(gin.ReleaseMode)
|
||||||
}
|
}
|
||||||
// Initialize SQL Database
|
// Initialize SQL Database
|
||||||
|
defer service.ShutdownWSHubs()
|
||||||
|
|
||||||
err := model.InitDB()
|
err := model.InitDB()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
slog.Error("initialize database failed", "error", err)
|
slog.Error("initialize database failed", "error", err)
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package service
|
|||||||
import (
|
import (
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
type WSMessage struct {
|
type WSMessage struct {
|
||||||
@@ -65,13 +66,58 @@ type WSHub struct {
|
|||||||
name string
|
name string
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
clients map[string]*WSClient
|
clients map[string]*WSClient
|
||||||
|
done chan struct{}
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewWSHub(name string) *WSHub {
|
func NewWSHub(name string) *WSHub {
|
||||||
return &WSHub{
|
h := &WSHub{
|
||||||
name: name,
|
name: name,
|
||||||
clients: make(map[string]*WSClient),
|
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 {
|
func (h *WSHub) Register(id string) *WSClient {
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package wsclient
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
@@ -25,6 +26,17 @@ type Client struct {
|
|||||||
cfg Config
|
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 {
|
type Connection struct {
|
||||||
Conn *websocket.Conn
|
Conn *websocket.Conn
|
||||||
URL string
|
URL string
|
||||||
@@ -155,6 +167,53 @@ func websocketReadTimeout(requestTimeout time.Duration) time.Duration {
|
|||||||
return timeout
|
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 {
|
func (conn *Connection) Close() error {
|
||||||
if conn == nil || conn.Conn == nil {
|
if conn == nil || conn.Conn == nil {
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -49,31 +49,31 @@ func (r *Runner) Run(ctx context.Context) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) {
|
type flaredWSHandler struct {
|
||||||
for {
|
runner *Runner
|
||||||
select {
|
}
|
||||||
case <-ctx.Done():
|
|
||||||
return
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
|
|
||||||
msg, err := conn.Receive()
|
func (h *flaredWSHandler) OnConnect(ctx context.Context) error {
|
||||||
if err != nil {
|
return nil
|
||||||
slog.Error("flared ws receive failed", "error", err)
|
}
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
switch msg.Type {
|
func (h *flaredWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
|
||||||
case "ping":
|
switch msg.Type {
|
||||||
_ = conn.SendPong()
|
case "active_config":
|
||||||
case "active_config":
|
slog.Info("received config update notification from server")
|
||||||
// Server notifies there is a new config available
|
h.runner.SyncService.Trigger()
|
||||||
slog.Info("received config update notification from server")
|
default:
|
||||||
r.SyncService.Trigger()
|
slog.Debug("ignored unknown ws message type", "type", msg.Type)
|
||||||
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) {
|
func (r *Runner) sleepContext(ctx context.Context, d time.Duration) {
|
||||||
|
|||||||
@@ -9,6 +9,9 @@ import (
|
|||||||
shared "openflare/utils/wsclient"
|
shared "openflare/utils/wsclient"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type WSMessage = shared.WSMessage
|
||||||
|
type MessageHandler = shared.MessageHandler
|
||||||
|
|
||||||
type Client struct {
|
type Client struct {
|
||||||
sharedClient *shared.Client
|
sharedClient *shared.Client
|
||||||
}
|
}
|
||||||
@@ -41,6 +44,10 @@ func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
|||||||
return &Connection{sharedConn: conn}, nil
|
return &Connection{sharedConn: conn}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (conn *Connection) SendPing() error {
|
||||||
|
return conn.sharedConn.SendMessage("ping", nil)
|
||||||
|
}
|
||||||
|
|
||||||
func (conn *Connection) SendPong() error {
|
func (conn *Connection) SendPong() error {
|
||||||
return conn.sharedConn.SendMessage("pong", nil)
|
return conn.sharedConn.SendMessage("pong", nil)
|
||||||
}
|
}
|
||||||
@@ -59,6 +66,10 @@ func (conn *Connection) Receive() (service.WSMessage, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error {
|
||||||
|
return conn.sharedConn.RunReceiveLoop(ctx, handler)
|
||||||
|
}
|
||||||
|
|
||||||
func (conn *Connection) Close() error {
|
func (conn *Connection) Close() error {
|
||||||
return conn.sharedConn.Close()
|
return conn.sharedConn.Close()
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user