package handler import ( "database/sql" "errors" "fmt" "net" "net/http" "strconv" "strings" "time" "go-backend/internal/ws" ) var errForwardNotFound = errors.New("forward not found") type forwardRecord struct { ID int64 UserID int64 UserName string Name string TunnelID int64 RemoteAddr string Strategy string Status int } type tunnelRecord struct { ID int64 Type int Status int Flow int64 TrafficRatio float64 } type forwardPortRecord struct { NodeID int64 Port int } type nodeRecord struct { ID int64 Name string ServerIP string TCPListenAddr string UDPListenAddr string InterfaceName string } type chainNodeRecord struct { ChainType int Inx int64 NodeID int64 NodeName string } func (h *Handler) resolveForwardAccess(r *http.Request, forwardID int64) (*forwardRecord, int64, int, error) { userID, roleID, err := userRoleFromRequest(r) if err != nil { return nil, 0, 0, err } forward, err := h.ensureForwardAccessByActor(userID, roleID, forwardID) if err != nil { return nil, userID, roleID, err } return forward, userID, roleID, nil } func (h *Handler) ensureForwardAccessByActor(actorUserID int64, actorRole int, forwardID int64) (*forwardRecord, error) { forward, err := h.getForwardRecord(forwardID) if err != nil { return nil, err } if actorRole != 0 && forward.UserID != actorUserID { return nil, errForwardNotFound } return forward, nil } func (h *Handler) ensureTunnelPermission(userID int64, roleID int, tunnelID int64) error { if roleID == 0 { return nil } var count int err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? AND status = 1`, userID, tunnelID).Scan(&count) if err != nil { return err } if count <= 0 { return errors.New("你没有该隧道的权限") } return nil } func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) { row := h.repo.DB().QueryRow(` SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status FROM forward WHERE id = ? LIMIT 1 `, forwardID) var fr forwardRecord err := row.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, errForwardNotFound } return nil, err } if strings.TrimSpace(fr.Strategy) == "" { fr.Strategy = "fifo" } return &fr, nil } func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) { row := h.repo.DB().QueryRow(`SELECT id, type, status, flow, traffic_ratio FROM tunnel WHERE id = ? LIMIT 1`, tunnelID) var tr tunnelRecord err := row.Scan(&tr.ID, &tr.Type, &tr.Status, &tr.Flow, &tr.TrafficRatio) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, errors.New("隧道不存在") } return nil, err } if tr.Flow <= 0 { tr.Flow = 1 } if tr.TrafficRatio <= 0 { tr.TrafficRatio = 1 } return &tr, nil } func (h *Handler) listForwardsByTunnel(tunnelID int64) ([]forwardRecord, error) { rows, err := h.repo.DB().Query(` SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status FROM forward WHERE tunnel_id = ? ORDER BY id ASC `, tunnelID) if err != nil { return nil, err } defer rows.Close() result := make([]forwardRecord, 0) for rows.Next() { var fr forwardRecord if err := rows.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status); err != nil { return nil, err } if strings.TrimSpace(fr.Strategy) == "" { fr.Strategy = "fifo" } result = append(result, fr) } if err := rows.Err(); err != nil { return nil, err } return result, nil } func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error) { rows, err := h.repo.DB().Query(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID) if err != nil { return nil, err } defer rows.Close() result := make([]forwardPortRecord, 0) for rows.Next() { var item forwardPortRecord if err := rows.Scan(&item.NodeID, &item.Port); err != nil { return nil, err } result = append(result, item) } if err := rows.Err(); err != nil { return nil, err } return result, nil } func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) { row := h.repo.DB().QueryRow(` SELECT id, name, server_ip, tcp_listen_addr, udp_listen_addr, interface_name FROM node WHERE id = ? LIMIT 1 `, nodeID) var n nodeRecord var tcpListen sql.NullString var udpListen sql.NullString var iface sql.NullString err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &tcpListen, &udpListen, &iface) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, errors.New("节点不存在") } return nil, err } n.TCPListenAddr = strings.TrimSpace(tcpListen.String) n.UDPListenAddr = strings.TrimSpace(udpListen.String) n.InterfaceName = strings.TrimSpace(iface.String) if n.TCPListenAddr == "" { n.TCPListenAddr = "[::]" } if n.UDPListenAddr == "" { n.UDPListenAddr = "[::]" } if strings.TrimSpace(n.Name) == "" { n.Name = fmt.Sprintf("node_%d", n.ID) } return &n, nil } func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int, error) { row := h.repo.DB().QueryRow(` SELECT ut.id, sl.speed FROM user_tunnel ut LEFT JOIN speed_limit sl ON sl.id = ut.speed_id WHERE ut.user_id = ? AND ut.tunnel_id = ? LIMIT 1 `, userID, tunnelID) var userTunnelID int64 var speed sql.NullInt64 err := row.Scan(&userTunnelID, &speed) if err != nil { if errors.Is(err, sql.ErrNoRows) { return 0, nil, nil } return 0, nil, err } if !speed.Valid || speed.Int64 <= 0 { return userTunnelID, nil, nil } v := int(speed.Int64) return userTunnelID, &v, nil } func (h *Handler) syncForwardServices(forward *forwardRecord, method string, allowFallbackAdd bool) error { if h == nil || forward == nil { return errors.New("invalid forward sync context") } tunnel, err := h.getTunnelRecord(forward.TunnelID) if err != nil { return err } ports, err := h.listForwardPorts(forward.ID) if err != nil { return err } if len(ports) == 0 { return errors.New("转发入口端口不存在") } userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) if err != nil { return err } serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID) for _, fp := range ports { node, err := h.getNodeRecord(fp.NodeID) if err != nil { return err } services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiter) _, err = h.sendNodeCommand(node.ID, method, services, true, false) if err != nil && allowFallbackAdd && method == "UpdateService" { _, err = h.sendNodeCommand(node.ID, "AddService", services, true, false) } if err != nil { return fmt.Errorf("节点 %s 下发失败: %w", node.Name, err) } } return nil } func (h *Handler) controlForwardServices(forward *forwardRecord, commandType string, tolerateNotFound bool) error { if h == nil || forward == nil { return errors.New("invalid forward control context") } ports, err := h.listForwardPorts(forward.ID) if err != nil { return err } if len(ports) == 0 { return nil } userTunnelID, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) if err != nil { return err } base := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID) payload := map[string]interface{}{ "services": []string{base, base + "_tcp", base + "_udp"}, } seen := map[int64]struct{}{} for _, fp := range ports { if _, ok := seen[fp.NodeID]; ok { continue } seen[fp.NodeID] = struct{}{} _, err := h.sendNodeCommand(fp.NodeID, commandType, payload, false, tolerateNotFound) if err != nil { return err } } return nil } func (h *Handler) applyNodeProtocolChange(nodeID int64, httpVal, tlsVal, socksVal int) error { _, err := h.sendNodeCommand(nodeID, "SetProtocol", map[string]interface{}{ "http": httpVal, "tls": tlsVal, "socks": socksVal, }, false, false) return err } func (h *Handler) sendNodeCommand(nodeID int64, commandType string, data interface{}, tolerateExists bool, tolerateNotFound bool) (ws.CommandResult, error) { result, err := h.wsServer.SendCommand(nodeID, commandType, data, 12*time.Second) if err == nil { return result, nil } msg := strings.ToLower(strings.TrimSpace(err.Error())) if tolerateExists { if strings.Contains(msg, "exists") || strings.Contains(msg, "already") || strings.Contains(msg, "已存在") { return result, nil } } if tolerateNotFound { if strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") { return result, nil } } return result, err } func (h *Handler) diagnoseForwardRuntime(forward *forwardRecord) (map[string]interface{}, error) { if forward == nil { return nil, errForwardNotFound } targets := splitRemoteTargets(forward.RemoteAddr) if len(targets) == 0 { return nil, errors.New("目标地址不能为空") } ip, port, err := parseTargetAddress(targets[0]) if err != nil { return nil, errors.New("目标地址格式错误") } ports, err := h.listForwardPorts(forward.ID) if err != nil { return nil, err } if len(ports) == 0 { return nil, errors.New("转发入口端口不存在") } results := make([]map[string]interface{}, 0, len(ports)) seen := map[int64]struct{}{} for _, fp := range ports { if _, ok := seen[fp.NodeID]; ok { continue } seen[fp.NodeID] = struct{}{} node, nodeErr := h.getNodeRecord(fp.NodeID) nodeName := fmt.Sprintf("node_%d", fp.NodeID) if nodeErr == nil { nodeName = node.Name } resultItem := map[string]interface{}{ "nodeName": nodeName, "nodeId": strconv.FormatInt(fp.NodeID, 10), "targetIp": ip, "targetPort": port, "averageTime": 0, "packetLoss": 100, } if nodeErr != nil { resultItem["success"] = false resultItem["description"] = "节点信息读取失败" resultItem["errorMessage"] = nodeErr.Error() results = append(results, resultItem) continue } pingData, pingErr := h.tcpPingViaNode(fp.NodeID, ip, port) if pingErr != nil { resultItem["success"] = false resultItem["description"] = "诊断失败" resultItem["errorMessage"] = pingErr.Error() } else { resultItem["success"] = asBool(pingData["success"], false) resultItem["description"] = "诊断完成" resultItem["averageTime"] = asFloat(pingData["averageTime"], 0) resultItem["packetLoss"] = asFloat(pingData["packetLoss"], 100) if msg := asString(pingData["errorMessage"]); msg != "" { resultItem["errorMessage"] = msg } } results = append(results, resultItem) } payload := map[string]interface{}{ "forwardName": forward.Name, "timestamp": time.Now().UnixMilli(), "results": results, } return payload, nil } func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{}, error) { tunnel, err := h.getTunnelRecord(tunnelID) if err != nil { return nil, err } var tunnelName string if err := h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName); err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, errors.New("隧道不存在") } return nil, err } chainRows, err := h.listChainNodesForTunnel(tunnelID) if err != nil { return nil, err } if len(chainRows) == 0 { return nil, errors.New("隧道配置不完整") } results := make([]map[string]interface{}, 0, len(chainRows)) seen := map[int64]struct{}{} for _, row := range chainRows { if _, ok := seen[row.NodeID]; ok { continue } seen[row.NodeID] = struct{}{} desc := fmt.Sprintf("节点(%s)->外网", row.NodeName) switch row.ChainType { case 1: desc = fmt.Sprintf("入口(%s)->外网", row.NodeName) case 2: desc = fmt.Sprintf("中继(%s)->外网", row.NodeName) case 3: desc = fmt.Sprintf("出口(%s)->外网", row.NodeName) } resultItem := map[string]interface{}{ "nodeName": row.NodeName, "nodeId": strconv.FormatInt(row.NodeID, 10), "targetIp": "www.google.com", "targetPort": 443, "averageTime": 0, "packetLoss": 100, "fromChainType": row.ChainType, "description": desc, } if row.ChainType == 2 { resultItem["fromInx"] = row.Inx } pingData, pingErr := h.tcpPingViaNode(row.NodeID, "www.google.com", 443) if pingErr != nil { resultItem["success"] = false resultItem["description"] = desc + " 连通性检查失败" resultItem["errorMessage"] = pingErr.Error() } else { resultItem["success"] = asBool(pingData["success"], false) resultItem["description"] = desc + " 连通性检查完成" resultItem["averageTime"] = asFloat(pingData["averageTime"], 0) resultItem["packetLoss"] = asFloat(pingData["packetLoss"], 100) if msg := asString(pingData["errorMessage"]); msg != "" { resultItem["errorMessage"] = msg } } results = append(results, resultItem) } payload := map[string]interface{}{ "tunnelName": tunnelName, "tunnelType": map[bool]string{true: "端口转发", false: "隧道转发"}[tunnel.Type == 1], "timestamp": time.Now().UnixMilli(), "results": results, } return payload, nil } func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) { rows, err := h.repo.DB().Query(` SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, n.name FROM chain_tunnel ct LEFT JOIN node n ON n.id = ct.node_id WHERE ct.tunnel_id = ? ORDER BY ct.chain_type ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC `, tunnelID) if err != nil { return nil, err } defer rows.Close() result := make([]chainNodeRecord, 0) for rows.Next() { var item chainNodeRecord var name sql.NullString if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &name); err != nil { return nil, err } if strings.TrimSpace(name.String) == "" { item.NodeName = fmt.Sprintf("node_%d", item.NodeID) } else { item.NodeName = name.String } result = append(result, item) } if err := rows.Err(); err != nil { return nil, err } return result, nil } func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int) (map[string]interface{}, error) { res, err := h.sendNodeCommand(nodeID, "TcpPing", map[string]interface{}{ "ip": ip, "port": port, "count": 4, "timeout": 5000, }, false, false) if err != nil { return nil, err } if res.Data == nil { return nil, errors.New("节点未返回诊断数据") } return res.Data, nil } func splitRemoteTargets(remoteAddr string) []string { parts := strings.Split(remoteAddr, ",") out := make([]string, 0, len(parts)) for _, part := range parts { part = strings.TrimSpace(part) if part == "" { continue } out = append(out, processServerAddress(part)) } return out } func parseTargetAddress(addr string) (string, int, error) { addr = strings.TrimSpace(addr) if addr == "" { return "", 0, errors.New("empty address") } host, portStr, err := net.SplitHostPort(addr) if err != nil { idx := strings.LastIndex(addr, ":") if idx <= 0 || idx >= len(addr)-1 { return "", 0, err } host = strings.TrimSpace(addr[:idx]) portStr = strings.TrimSpace(addr[idx+1:]) } port, err := strconv.Atoi(strings.TrimSpace(portStr)) if err != nil || port <= 0 || port > 65535 { return "", 0, errors.New("invalid port") } host = strings.Trim(strings.TrimSpace(host), "[]") if host == "" { return "", 0, errors.New("invalid host") } return host, port, nil } func buildForwardServiceBase(forwardID, userID, userTunnelID int64) string { return fmt.Sprintf("%d_%d_%d", forwardID, userID, userTunnelID) } func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiter *int) []map[string]interface{} { protocols := []string{"tcp", "udp"} services := make([]map[string]interface{}, 0, 2) targets := splitRemoteTargets(forward.RemoteAddr) strategy := strings.TrimSpace(forward.Strategy) if strategy == "" { strategy = "fifo" } for _, protocol := range protocols { listenerAddr := node.TCPListenAddr if protocol == "udp" { listenerAddr = node.UDPListenAddr } service := map[string]interface{}{ "name": fmt.Sprintf("%s_%s", baseName, protocol), "addr": fmt.Sprintf("%s:%d", listenerAddr, port), "handler": map[string]interface{}{ "type": protocol, }, "listener": map[string]interface{}{ "type": protocol, }, "forwarder": map[string]interface{}{ "nodes": buildForwarderNodes(targets), "selector": map[string]interface{}{ "strategy": strategy, "maxFails": 1, "failTimeout": "600s", }, }, } if protocol == "udp" { service["listener"].(map[string]interface{})["metadata"] = map[string]interface{}{"keepAlive": true} } if tunnel != nil && tunnel.Type == 2 { service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID) } if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" { service["metadata"] = map[string]interface{}{"interface": node.InterfaceName} } if limiter != nil && *limiter > 0 { service["limiter"] = strconv.Itoa(*limiter) } services = append(services, service) } return services } func buildForwarderNodes(targets []string) []map[string]interface{} { nodes := make([]map[string]interface{}, 0, len(targets)) for i, addr := range targets { nodes = append(nodes, map[string]interface{}{ "name": fmt.Sprintf("node_%d", i+1), "addr": addr, }) } return nodes } func processServerAddress(serverAddr string) string { serverAddr = strings.TrimSpace(serverAddr) if serverAddr == "" { return serverAddr } if strings.HasPrefix(serverAddr, "[") { return serverAddr } idx := strings.LastIndex(serverAddr, ":") if idx < 0 { if looksLikeIPv6(serverAddr) { return "[" + serverAddr + "]" } return serverAddr } host := strings.TrimSpace(serverAddr[:idx]) port := strings.TrimSpace(serverAddr[idx+1:]) if host == "" || port == "" { return serverAddr } if looksLikeIPv6(host) { return "[" + host + "]:" + port } return serverAddr } func looksLikeIPv6(address string) bool { return strings.Count(address, ":") >= 2 } func asBool(v interface{}, def bool) bool { s := strings.TrimSpace(strings.ToLower(asString(v))) if s == "" { return def } switch s { case "1", "t", "true", "yes", "y": return true case "0", "f", "false", "no", "n": return false default: return def } }