package handler import ( "encoding/json" "net/http" "time" "git.zkcoi.com/zkcoi/meshray/internal/ctr" "git.zkcoi.com/zkcoi/meshray/internal/model" "git.zkcoi.com/zkcoi/meshray/internal/store/sqlite" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" "go.uber.org/zap" ) var upgrader = websocket.Upgrader{ // ✅ 修复:限制 WebSocket 跨域,只允许特定来源 CheckOrigin: func(r *http.Request) bool { origin := r.Header.Get("Origin") allowedOrigins := map[string]bool{ "http://localhost:9531": true, "http://127.0.0.1:9531": true, } return allowedOrigins[origin] }, } // WSMessage 与前端定义一致的结构 type WSMessage struct { Channel string `json:"channel"` Type string `json:"type"` Payload interface{} `json:"payload"` } // WSHandler 处理 WebSocket 连接 type WSHandler struct { ctrClient *ctr.Ctr logger *zap.Logger store *sqlite.Store } func NewWSHandler(ctrClient *ctr.Ctr, store *sqlite.Store, logger *zap.Logger) *WSHandler { return &WSHandler{ ctrClient: ctrClient, logger: logger, store: store, } } func (h *WSHandler) ServeWS(c *gin.Context) { ws, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { h.logger.Error("升级 WebSocket 失败", zap.Error(err)) return } defer ws.Close() h.logger.Info("客户端已连接 WebSocket") // 这里暂时用简单轮询的方案,每秒向终端推送 core metrics 和 peering 状态 // 在生产级应用中,可以使用 EventBus 订阅发布模式 ticker := time.NewTicker(1 * time.Second) defer ticker.Stop() // 启动一个 goroutine 读取客户端发来的控制指令 (ping/subscribe) go func() { for { _, msg, err := ws.ReadMessage() if err != nil { h.logger.Debug("WebSocket 客户端断开", zap.Error(err)) break } // 简单响应 ping var req map[string]interface{} if err := json.Unmarshal(msg, &req); err == nil { if req["type"] == "ping" { _ = ws.WriteJSON(map[string]interface{}{"type": "pong", "timestamp": time.Now().Unix()}) } } } }() for { select { case <-ticker.C: // 1. 获取所有组网 var networks []model.Network if err := h.store.DB().Find(&networks).Error; err != nil { continue } // 聚合数据 var totalBytesSent, totalBytesRecv uint64 var totalFallbacks int64 var totalActiveConnections int64 // 收集链路分布 (Layer -> Count) layerDistribution := make(map[string]int) // wg 状态聚合(设备是否在运行) wgRunningCount := 0 // 2. 遍历所有组网聚合 Ctr Status for _, net := range networks { status, err := h.ctrClient.GetStatus(net.ID) if err != nil || status == nil { continue } if status.WGStatus != nil && status.WGStatus.Running { wgRunningCount++ } if status.CoreStatus != nil { totalBytesSent += status.CoreStatus.BytesSent totalBytesRecv += status.CoreStatus.BytesReceived totalFallbacks += status.CoreStatus.StrategyFallbacks totalActiveConnections += status.CoreStatus.ActiveConnections for _, peerStatus := range status.CoreStatus.Peers { if peerStatus.Connected && peerStatus.Layer != "" { layerDistribution[peerStatus.Layer]++ } else if peerStatus.Connected { layerDistribution["Native"]++ } } } } // 3. 构建 payload 推送 msg := WSMessage{ Channel: "metrics", Type: "update", Payload: map[string]interface{}{ "summary": map[string]interface{}{ "networks_count": len(networks), "wg_running": wgRunningCount, }, "core_status": map[string]interface{}{ "bytes_sent": totalBytesSent, "bytes_received": totalBytesRecv, "strategy_fallbacks": totalFallbacks, "active_connections": totalActiveConnections, "layer_distribution": layerDistribution, }, }, } if err := ws.WriteJSON(msg); err != nil { return } case <-c.Request.Context().Done(): return } } }