Files
flvx/go-backend/internal/ws/server.go
T
2026-05-15 23:37:08 +08:00

757 lines
18 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 nodeSession struct {
nodeID int64
secret string
conn *websocket.Conn
crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建
}
type adminSession struct {
userID int64
claims auth.Claims
conn *websocket.Conn
}
type monitorSession struct {
userID int64
claims auth.Claims
conn *websocket.Conn
}
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)
getUserAuthState func(userID int64) (*auth.UserAuthState, error)
mu sync.RWMutex
admins map[*adminSession]struct{}
monitors map[*monitorSession]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[*adminSession]struct{}),
monitors: make(map[*monitorSession]struct{}),
nodes: make(map[int64]*nodeSession),
byConn: make(map[*websocket.Conn]*nodeSession),
pending: make(map[string]pendingRequest),
}
}
func (s *Server) SetUserAuthStateLookup(fn func(userID int64) (*auth.UserAuthState, error)) {
if s == nil {
return
}
s.mu.Lock()
s.getUserAuthState = fn
s.mu.Unlock()
}
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" {
claims, ok := auth.ValidateToken(secret, s.jwtSecret)
if !ok {
http.Error(w, "forbidden", http.StatusForbidden)
return
}
userID, err := strconv.ParseInt(claims.Sub, 10, 64)
if err != nil {
http.Error(w, "forbidden", http.StatusForbidden)
return
}
if claims.RoleID == 0 {
if !s.validateAdminSession(userID, claims) {
http.Error(w, "forbidden", http.StatusForbidden)
return
}
s.handleAdmin(w, r, userID, claims)
return
}
if !s.validateMonitorSession(userID, claims) {
http.Error(w, "forbidden", http.StatusForbidden)
return
}
s.handleMonitor(w, r, userID, claims)
return
}
http.Error(w, "bad request", http.StatusBadRequest)
}
func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request, userID int64, claims auth.Claims) {
conn, err := s.upgrader.Upgrade(w, r, nil)
if err != nil {
return
}
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
})
done := make(chan struct{})
session := &adminSession{userID: userID, claims: claims, conn: conn}
go startKeepalive(conn, done, func() bool {
return s.validateAdminSession(session.userID, session.claims)
})
s.mu.Lock()
s.admins[session] = struct{}{}
s.mu.Unlock()
defer func() {
close(done)
s.mu.Lock()
delete(s.admins, session)
s.mu.Unlock()
_ = conn.Close()
}()
for {
if _, _, err := conn.ReadMessage(); err != nil {
return
}
}
}
func (s *Server) handleMonitor(w http.ResponseWriter, r *http.Request, userID int64, claims auth.Claims) {
conn, err := s.upgrader.Upgrade(w, r, nil)
if err != nil {
return
}
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
})
done := make(chan struct{})
session := &monitorSession{userID: userID, claims: claims, conn: conn}
go startKeepalive(conn, done, func() bool {
return s.validateMonitorSession(session.userID, session.claims)
})
s.mu.Lock()
s.monitors[session] = struct{}{}
s.mu.Unlock()
defer func() {
close(done)
s.mu.Lock()
delete(s.monitors, session)
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
}
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
})
done := make(chan struct{})
go startKeepalive(conn, done, nil)
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.Close()
delete(s.byConn, old.conn)
}
// 初始化 AES 加密器并缓存(仅创建一次)
var nodeCrypto *security.AESCrypto
if strings.TrimSpace(secret) != "" {
nodeCrypto, _ = security.NewAESCrypto(secret)
}
ns := &nodeSession{nodeID: nodeID, secret: secret, conn: conn, 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 {
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 {
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.SetWriteDeadline(time.Now().Add(wsWriteWait))
err = ns.conn.WriteMessage(websocket.TextMessage, messageData)
_ = ns.conn.SetWriteDeadline(time.Time{})
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.broadcastToRealtime(string(raw))
}
func (s *Server) broadcastInfo(nodeID int64, data string) {
payload := broadcastMessage{ID: nodeID, Type: "info", Data: data}
raw, _ := json.Marshal(payload)
s.broadcastToRealtime(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.broadcastToRealtime(string(raw))
}
func (s *Server) broadcastToRealtime(message string) {
s.mu.RLock()
admins := make([]*adminSession, 0, len(s.admins))
for c := range s.admins {
admins = append(admins, c)
}
monitors := make([]*monitorSession, 0, len(s.monitors))
for c := range s.monitors {
monitors = append(monitors, c)
}
s.mu.RUnlock()
for _, c := range admins {
if c == nil || c.conn == nil {
continue
}
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
_ = c.conn.SetWriteDeadline(time.Time{})
if err != nil {
log.Printf("websocket broadcast failed: %v", err)
}
}
for _, c := range monitors {
if c == nil || c.conn == nil {
continue
}
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
_ = c.conn.SetWriteDeadline(time.Time{})
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 (s *Server) validateAdminSession(userID int64, claims auth.Claims) bool {
if s == nil {
return false
}
if claims.Exp <= time.Now().Unix() {
return false
}
if s.getUserAuthState == nil {
return true
}
state, err := s.getUserAuthState(userID)
if err != nil || state == nil || state.Status != 1 || state.RoleID != claims.RoleID || claims.IatMs <= state.PasswordChangedAt {
return false
}
return true
}
func (s *Server) validateMonitorSession(userID int64, claims auth.Claims) bool {
if s == nil {
return false
}
if claims.Exp <= time.Now().Unix() {
return false
}
if s.getUserAuthState != nil {
state, err := s.getUserAuthState(userID)
if err != nil || state == nil || state.Status != 1 || state.RoleID != claims.RoleID || claims.IatMs <= state.PasswordChangedAt {
return false
}
}
if s.repo == nil {
return false
}
allowed, err := s.repo.HasMonitorPermission(userID)
if err != nil || !allowed {
return false
}
return true
}
func startKeepalive(conn *websocket.Conn, done <-chan struct{}, validate func() bool) {
if conn == nil {
return
}
ticker := time.NewTicker(wsPingPeriod)
defer ticker.Stop()
for {
select {
case <-done:
return
case <-ticker.C:
if validate != nil && !validate() {
_ = conn.Close()
return
}
_ = conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
err := conn.WriteMessage(websocket.PingMessage, nil)
_ = conn.SetWriteDeadline(time.Time{})
if err != nil {
_ = conn.Close()
return
}
}
}
}