mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
629 lines
14 KiB
Go
629 lines
14 KiB
Go
package ws
|
|
|
|
import (
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"log"
|
|
"net/http"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/gorilla/websocket"
|
|
|
|
"go-backend/internal/auth"
|
|
"go-backend/internal/security"
|
|
"go-backend/internal/store/repo"
|
|
)
|
|
|
|
type encryptedMessage struct {
|
|
Encrypted bool `json:"encrypted"`
|
|
Data string `json:"data"`
|
|
Timestamp int64 `json:"timestamp"`
|
|
}
|
|
|
|
type broadcastMessage struct {
|
|
ID int64 `json:"id"`
|
|
Type string `json:"type"`
|
|
Data string `json:"data"`
|
|
}
|
|
|
|
type connWrap struct {
|
|
conn *websocket.Conn
|
|
mu sync.Mutex
|
|
}
|
|
|
|
type nodeSession struct {
|
|
nodeID int64
|
|
secret string
|
|
conn *connWrap
|
|
crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建
|
|
}
|
|
|
|
type commandResponse struct {
|
|
Type string `json:"type"`
|
|
Success bool `json:"success"`
|
|
Message string `json:"message"`
|
|
Data json.RawMessage `json:"data,omitempty"`
|
|
RequestID string `json:"requestId,omitempty"`
|
|
}
|
|
|
|
type pendingRequest struct {
|
|
nodeID int64
|
|
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"`
|
|
Message string `json:"message"`
|
|
Data map[string]interface{} `json:"data,omitempty"`
|
|
}
|
|
|
|
type Server struct {
|
|
repo *repo.Repository
|
|
jwtSecret string
|
|
upgrader websocket.Upgrader
|
|
onNodeOnline func(nodeID int64)
|
|
onNodeMetric func(nodeID int64, info SystemInfo)
|
|
|
|
mu sync.RWMutex
|
|
admins map[*connWrap]struct{}
|
|
nodes map[int64]*nodeSession
|
|
byConn map[*websocket.Conn]*nodeSession
|
|
pending map[string]pendingRequest
|
|
}
|
|
|
|
type SystemInfo struct {
|
|
Uptime uint64 `json:"uptime"`
|
|
BytesReceived uint64 `json:"bytes_received"`
|
|
BytesTransmitted uint64 `json:"bytes_transmitted"`
|
|
CPUUsage float64 `json:"cpu_usage"`
|
|
MemoryUsage float64 `json:"memory_usage"`
|
|
DiskUsage float64 `json:"disk_usage"`
|
|
Load1 float64 `json:"load1"`
|
|
Load5 float64 `json:"load5"`
|
|
Load15 float64 `json:"load15"`
|
|
TCPConns int64 `json:"tcp_conns"`
|
|
UDPConns int64 `json:"udp_conns"`
|
|
NetInSpeed int64 `json:"net_in_speed"`
|
|
NetOutSpeed int64 `json:"net_out_speed"`
|
|
}
|
|
|
|
func (s *Server) SetNodeOnlineHook(fn func(nodeID int64)) {
|
|
if s == nil {
|
|
return
|
|
}
|
|
s.mu.Lock()
|
|
s.onNodeOnline = fn
|
|
s.mu.Unlock()
|
|
}
|
|
|
|
func (s *Server) SetNodeMetricHook(fn func(nodeID int64, info SystemInfo)) {
|
|
if s == nil {
|
|
return
|
|
}
|
|
s.mu.Lock()
|
|
s.onNodeMetric = fn
|
|
s.mu.Unlock()
|
|
}
|
|
|
|
func NewServer(repo *repo.Repository, jwtSecret string) *Server {
|
|
return &Server{
|
|
repo: repo,
|
|
jwtSecret: jwtSecret,
|
|
upgrader: websocket.Upgrader{
|
|
CheckOrigin: func(r *http.Request) bool { return true },
|
|
},
|
|
admins: make(map[*connWrap]struct{}),
|
|
nodes: make(map[int64]*nodeSession),
|
|
byConn: make(map[*websocket.Conn]*nodeSession),
|
|
pending: make(map[string]pendingRequest),
|
|
}
|
|
}
|
|
|
|
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|
query := r.URL.Query()
|
|
typeVal := query.Get("type")
|
|
secret := query.Get("secret")
|
|
|
|
if typeVal == "1" {
|
|
node, err := s.repo.GetNodeBySecret(secret)
|
|
if err != nil || node == nil {
|
|
http.Error(w, "forbidden", http.StatusForbidden)
|
|
return
|
|
}
|
|
s.handleNode(w, r, node.ID, secret)
|
|
return
|
|
}
|
|
|
|
if typeVal == "0" {
|
|
if _, ok := auth.ValidateToken(secret, s.jwtSecret); !ok {
|
|
http.Error(w, "forbidden", http.StatusForbidden)
|
|
return
|
|
}
|
|
s.handleAdmin(w, r)
|
|
return
|
|
}
|
|
|
|
http.Error(w, "bad request", http.StatusBadRequest)
|
|
}
|
|
|
|
func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) {
|
|
conn, err := s.upgrader.Upgrade(w, r, nil)
|
|
if err != nil {
|
|
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()
|
|
_ = conn.Close()
|
|
}()
|
|
|
|
for {
|
|
if _, _, err := conn.ReadMessage(); err != nil {
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64, secret string) {
|
|
conn, err := s.upgrader.Upgrade(w, r, nil)
|
|
if err != nil {
|
|
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)
|
|
tlsVal := parseIntDefault(r.URL.Query().Get("tls"), 0)
|
|
socksVal := parseIntDefault(r.URL.Query().Get("socks"), 0)
|
|
|
|
s.mu.Lock()
|
|
if old, ok := s.nodes[nodeID]; ok {
|
|
_ = old.conn.conn.Close()
|
|
delete(s.byConn, old.conn.conn)
|
|
}
|
|
// 初始化 AES 加密器并缓存(仅创建一次)
|
|
var nodeCrypto *security.AESCrypto
|
|
if strings.TrimSpace(secret) != "" {
|
|
nodeCrypto, _ = security.NewAESCrypto(secret)
|
|
}
|
|
ns := &nodeSession{nodeID: nodeID, secret: secret, conn: cw, crypto: nodeCrypto}
|
|
s.nodes[nodeID] = ns
|
|
s.byConn[conn] = ns
|
|
s.mu.Unlock()
|
|
|
|
_ = s.repo.UpdateNodeOnline(nodeID, 1, version, httpVal, tlsVal, socksVal)
|
|
s.broadcastStatus(nodeID, 1)
|
|
|
|
s.mu.RLock()
|
|
onlineHook := s.onNodeOnline
|
|
s.mu.RUnlock()
|
|
if onlineHook != nil {
|
|
go onlineHook(nodeID)
|
|
}
|
|
|
|
defer func() {
|
|
close(done)
|
|
needOfflineBroadcast := false
|
|
s.mu.Lock()
|
|
current, ok := s.nodes[nodeID]
|
|
if ok && current.conn.conn == conn {
|
|
delete(s.nodes, nodeID)
|
|
needOfflineBroadcast = true
|
|
}
|
|
delete(s.byConn, conn)
|
|
s.mu.Unlock()
|
|
if needOfflineBroadcast {
|
|
s.failPendingForNode(nodeID, "节点连接已断开")
|
|
_ = s.repo.UpdateNodeStatus(nodeID, 0)
|
|
s.broadcastStatus(nodeID, 0)
|
|
}
|
|
_ = conn.Close()
|
|
}()
|
|
|
|
for {
|
|
_, payload, err := conn.ReadMessage()
|
|
if err != nil {
|
|
return
|
|
}
|
|
|
|
msg := decryptIfNeeded(payload, ns.crypto, secret)
|
|
s.tryResolvePending(nodeID, msg)
|
|
|
|
var parsed struct {
|
|
Type string `json:"type"`
|
|
}
|
|
if json.Unmarshal([]byte(msg), &parsed) == nil && parsed.Type != "" {
|
|
switch parsed.Type {
|
|
case "metric":
|
|
// Agent 新版指标消息:{type:"metric", data:{...}}
|
|
var envelope struct {
|
|
Data json.RawMessage `json:"data"`
|
|
}
|
|
if err := json.Unmarshal([]byte(msg), &envelope); err == nil && len(envelope.Data) > 0 {
|
|
// 解析 SystemInfo 并调用 hook
|
|
var sysInfo SystemInfo
|
|
if json.Unmarshal(envelope.Data, &sysInfo) == nil {
|
|
s.mu.RLock()
|
|
onMetric := s.onNodeMetric
|
|
s.mu.RUnlock()
|
|
if onMetric != nil {
|
|
go onMetric(nodeID, sysInfo)
|
|
}
|
|
}
|
|
// 广播内层 data 给前端(保持平坦结构兼容性)
|
|
s.broadcastTyped(nodeID, "metric", string(envelope.Data))
|
|
}
|
|
continue
|
|
case "UpgradeProgress":
|
|
s.broadcastTyped(nodeID, "upgrade_progress", msg)
|
|
continue
|
|
default:
|
|
// Unknown typed messages still get broadcast so future
|
|
// agent message types are not silently lost.
|
|
s.broadcastInfo(nodeID, msg)
|
|
continue
|
|
}
|
|
}
|
|
|
|
// 兼容旧版 Agent:无 type 字段的系统信息消息
|
|
if looksLikeSystemInfoMessage(msg) {
|
|
var sysInfo SystemInfo
|
|
if err := json.Unmarshal([]byte(msg), &sysInfo); err == nil {
|
|
s.mu.RLock()
|
|
onMetric := s.onNodeMetric
|
|
s.mu.RUnlock()
|
|
if onMetric != nil {
|
|
go onMetric(nodeID, sysInfo)
|
|
}
|
|
s.broadcastTyped(nodeID, "metric", msg)
|
|
continue
|
|
}
|
|
}
|
|
|
|
s.broadcastInfo(nodeID, msg)
|
|
}
|
|
}
|
|
|
|
func looksLikeSystemInfoMessage(msg string) bool {
|
|
// Keep this as a cheap heuristic so that arbitrary JSON objects don't get
|
|
// misclassified as metrics (SystemInfo unmarshal would otherwise succeed with
|
|
// all-zero values).
|
|
if strings.TrimSpace(msg) == "" {
|
|
return false
|
|
}
|
|
if !strings.Contains(msg, "{") {
|
|
return false
|
|
}
|
|
|
|
keys := []string{
|
|
"\"uptime\"",
|
|
"\"cpu_usage\"",
|
|
"\"memory_usage\"",
|
|
"\"disk_usage\"",
|
|
"\"bytes_received\"",
|
|
"\"bytes_transmitted\"",
|
|
"\"net_in_speed\"",
|
|
"\"net_out_speed\"",
|
|
"\"tcp_conns\"",
|
|
"\"udp_conns\"",
|
|
"\"load1\"",
|
|
"\"load5\"",
|
|
"\"load15\"",
|
|
}
|
|
matched := 0
|
|
for _, k := range keys {
|
|
if strings.Contains(msg, k) {
|
|
matched++
|
|
if matched >= 3 {
|
|
return true
|
|
}
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, timeout time.Duration) (CommandResult, error) {
|
|
if s == nil {
|
|
return CommandResult{}, errors.New("server not initialized")
|
|
}
|
|
if strings.TrimSpace(cmdType) == "" {
|
|
return CommandResult{}, errors.New("command type is empty")
|
|
}
|
|
if timeout <= 0 {
|
|
timeout = 10 * time.Second
|
|
}
|
|
|
|
s.mu.RLock()
|
|
ns, ok := s.nodes[nodeID]
|
|
s.mu.RUnlock()
|
|
if !ok || ns == nil || ns.conn == nil || ns.conn.conn == nil {
|
|
return CommandResult{}, errors.New("节点不在线")
|
|
}
|
|
|
|
requestID := fmt.Sprintf("%d_%d", nodeID, time.Now().UnixNano())
|
|
ch := make(chan CommandResult, 1)
|
|
|
|
s.mu.Lock()
|
|
s.pending[requestID] = pendingRequest{nodeID: nodeID, ch: ch}
|
|
s.mu.Unlock()
|
|
|
|
cleanup := func() {
|
|
s.mu.Lock()
|
|
if p, exists := s.pending[requestID]; exists {
|
|
delete(s.pending, requestID)
|
|
close(p.ch)
|
|
}
|
|
s.mu.Unlock()
|
|
}
|
|
|
|
cmdPayload := map[string]interface{}{
|
|
"type": cmdType,
|
|
"data": data,
|
|
"requestId": requestID,
|
|
}
|
|
rawCmd, err := json.Marshal(cmdPayload)
|
|
if err != nil {
|
|
cleanup()
|
|
return CommandResult{}, err
|
|
}
|
|
|
|
messageData := rawCmd
|
|
if ns.crypto != nil {
|
|
encrypted, err := ns.crypto.Encrypt(rawCmd)
|
|
if err != nil {
|
|
cleanup()
|
|
return CommandResult{}, err
|
|
}
|
|
wrapper := map[string]interface{}{
|
|
"encrypted": true,
|
|
"data": encrypted,
|
|
"timestamp": time.Now().UnixMilli(),
|
|
}
|
|
messageData, err = json.Marshal(wrapper)
|
|
if err != nil {
|
|
cleanup()
|
|
return CommandResult{}, err
|
|
}
|
|
}
|
|
|
|
ns.conn.mu.Lock()
|
|
_ = ns.conn.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
|
err = ns.conn.conn.WriteMessage(websocket.TextMessage, messageData)
|
|
_ = ns.conn.conn.SetWriteDeadline(time.Time{})
|
|
ns.conn.mu.Unlock()
|
|
if err != nil {
|
|
cleanup()
|
|
return CommandResult{}, err
|
|
}
|
|
|
|
select {
|
|
case result, ok := <-ch:
|
|
if !ok {
|
|
return CommandResult{}, errors.New("命令通道已关闭")
|
|
}
|
|
if !result.Success {
|
|
if strings.TrimSpace(result.Message) == "" {
|
|
result.Message = "命令执行失败"
|
|
}
|
|
return result, errors.New(result.Message)
|
|
}
|
|
return result, nil
|
|
case <-time.After(timeout):
|
|
cleanup()
|
|
return CommandResult{}, errors.New("等待节点响应超时")
|
|
}
|
|
}
|
|
|
|
func (s *Server) tryResolvePending(nodeID int64, message string) {
|
|
if s == nil || strings.TrimSpace(message) == "" {
|
|
return
|
|
}
|
|
|
|
// 快速短路:指标消息永远不含 requestId,跳过完整 JSON 解析
|
|
if !strings.Contains(message, "\"requestId\"") {
|
|
return
|
|
}
|
|
|
|
var resp commandResponse
|
|
if err := json.Unmarshal([]byte(message), &resp); err != nil {
|
|
return
|
|
}
|
|
if strings.TrimSpace(resp.RequestID) == "" {
|
|
return
|
|
}
|
|
|
|
s.mu.Lock()
|
|
p, ok := s.pending[resp.RequestID]
|
|
if ok {
|
|
delete(s.pending, resp.RequestID)
|
|
}
|
|
s.mu.Unlock()
|
|
if !ok {
|
|
return
|
|
}
|
|
if p.nodeID != nodeID {
|
|
select {
|
|
case p.ch <- CommandResult{Type: resp.Type, Success: false, Message: "节点响应与请求不匹配"}:
|
|
default:
|
|
}
|
|
close(p.ch)
|
|
return
|
|
}
|
|
|
|
result := CommandResult{
|
|
Type: resp.Type,
|
|
Success: resp.Success,
|
|
Message: resp.Message,
|
|
}
|
|
if len(resp.Data) > 0 {
|
|
var data map[string]interface{}
|
|
if err := json.Unmarshal(resp.Data, &data); err == nil {
|
|
result.Data = data
|
|
}
|
|
}
|
|
|
|
select {
|
|
case p.ch <- result:
|
|
default:
|
|
}
|
|
close(p.ch)
|
|
}
|
|
|
|
func (s *Server) failPendingForNode(nodeID int64, message string) {
|
|
if s == nil {
|
|
return
|
|
}
|
|
|
|
type pair struct {
|
|
id string
|
|
pr pendingRequest
|
|
}
|
|
items := make([]pair, 0)
|
|
|
|
s.mu.Lock()
|
|
for id, pr := range s.pending {
|
|
if pr.nodeID != nodeID {
|
|
continue
|
|
}
|
|
items = append(items, pair{id: id, pr: pr})
|
|
delete(s.pending, id)
|
|
}
|
|
s.mu.Unlock()
|
|
|
|
for _, item := range items {
|
|
select {
|
|
case item.pr.ch <- CommandResult{Success: false, Message: message}:
|
|
default:
|
|
}
|
|
close(item.pr.ch)
|
|
}
|
|
}
|
|
|
|
func (s *Server) broadcastStatus(nodeID int64, status int) {
|
|
payload := map[string]interface{}{
|
|
"id": strconv.FormatInt(nodeID, 10),
|
|
"type": "status",
|
|
"data": status,
|
|
}
|
|
raw, _ := json.Marshal(payload)
|
|
s.broadcastToAdmins(string(raw))
|
|
}
|
|
|
|
func (s *Server) broadcastInfo(nodeID int64, data string) {
|
|
payload := broadcastMessage{ID: nodeID, Type: "info", Data: data}
|
|
raw, _ := json.Marshal(payload)
|
|
s.broadcastToAdmins(string(raw))
|
|
}
|
|
|
|
func (s *Server) broadcastTyped(nodeID int64, msgType string, data string) {
|
|
payload := broadcastMessage{ID: nodeID, Type: msgType, Data: data}
|
|
raw, _ := json.Marshal(payload)
|
|
s.broadcastToAdmins(string(raw))
|
|
}
|
|
|
|
func (s *Server) broadcastToAdmins(message string) {
|
|
s.mu.RLock()
|
|
admins := make([]*connWrap, 0, len(s.admins))
|
|
for c := range s.admins {
|
|
admins = append(admins, c)
|
|
}
|
|
s.mu.RUnlock()
|
|
|
|
for _, c := range admins {
|
|
c.mu.Lock()
|
|
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
|
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
|
|
_ = c.conn.SetWriteDeadline(time.Time{})
|
|
c.mu.Unlock()
|
|
if err != nil {
|
|
log.Printf("websocket broadcast failed: %v", err)
|
|
}
|
|
}
|
|
}
|
|
|
|
func decryptIfNeeded(payload []byte, crypto *security.AESCrypto, secret string) string {
|
|
text := string(payload)
|
|
var wrap encryptedMessage
|
|
if err := json.Unmarshal(payload, &wrap); err != nil || !wrap.Encrypted || strings.TrimSpace(wrap.Data) == "" {
|
|
return text
|
|
}
|
|
|
|
// 优先使用缓存的 crypto 实例
|
|
c := crypto
|
|
if c == nil && strings.TrimSpace(secret) != "" {
|
|
c, _ = security.NewAESCrypto(secret)
|
|
}
|
|
if c == nil {
|
|
return text
|
|
}
|
|
plain, err := c.Decrypt(wrap.Data)
|
|
if err != nil {
|
|
return text
|
|
}
|
|
return string(plain)
|
|
}
|
|
|
|
func parseIntDefault(v string, fallback int) int {
|
|
x, err := strconv.Atoi(v)
|
|
if err != nil {
|
|
return fallback
|
|
}
|
|
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.conn.SetWriteDeadline(time.Time{})
|
|
cw.mu.Unlock()
|
|
if err != nil {
|
|
_ = cw.conn.Close()
|
|
return
|
|
}
|
|
}
|
|
}
|
|
}
|