diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go new file mode 100644 index 0000000..5830c14 --- /dev/null +++ b/go-backend/internal/http/handler/control_plane.go @@ -0,0 +1,685 @@ +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 +} + +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 FROM tunnel WHERE id = ? LIMIT 1`, tunnelID) + var tr tunnelRecord + err := row.Scan(&tr.ID, &tr.Type, &tr.Status) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, errors.New("隧道不存在") + } + return nil, err + } + 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 + "_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 + } +} diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 631554b..391dd0c 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -3,6 +3,7 @@ package handler import ( "database/sql" "encoding/json" + "fmt" "io" "net/http" "sort" @@ -63,26 +64,80 @@ func (h *Handler) WebSocketHandler() http.Handler { func (h *Handler) Register(mux *http.ServeMux) { mux.HandleFunc("/api/v1/user/login", h.login) mux.HandleFunc("/api/v1/user/list", h.userList) + mux.HandleFunc("/api/v1/user/create", h.userCreate) + mux.HandleFunc("/api/v1/user/update", h.userUpdate) + mux.HandleFunc("/api/v1/user/delete", h.userDelete) + mux.HandleFunc("/api/v1/user/reset", h.userResetFlow) mux.HandleFunc("/api/v1/config/get", h.getConfigByName) mux.HandleFunc("/api/v1/config/list", h.getConfigs) mux.HandleFunc("/api/v1/config/update", h.updateConfigs) mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig) mux.HandleFunc("/api/v1/captcha/check", h.checkCaptcha) + mux.HandleFunc("/api/v1/captcha/generate", h.captchaGenerate) + mux.HandleFunc("/api/v1/captcha/verify", h.captchaVerify) mux.HandleFunc("/api/v1/user/package", h.userPackage) mux.HandleFunc("/api/v1/user/updatePassword", h.updatePassword) mux.HandleFunc("/api/v1/node/list", h.nodeList) + mux.HandleFunc("/api/v1/node/create", h.nodeCreate) + mux.HandleFunc("/api/v1/node/update", h.nodeUpdate) + mux.HandleFunc("/api/v1/node/delete", h.nodeDelete) + mux.HandleFunc("/api/v1/node/install", h.nodeInstall) + mux.HandleFunc("/api/v1/node/update-order", h.nodeUpdateOrder) + mux.HandleFunc("/api/v1/node/batch-delete", h.nodeBatchDelete) + mux.HandleFunc("/api/v1/node/check-status", h.nodeCheckStatus) mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList) + mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate) + mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet) + mux.HandleFunc("/api/v1/tunnel/update", h.tunnelUpdate) + mux.HandleFunc("/api/v1/tunnel/delete", h.tunnelDelete) + mux.HandleFunc("/api/v1/tunnel/diagnose", h.tunnelDiagnose) + mux.HandleFunc("/api/v1/tunnel/update-order", h.tunnelUpdateOrder) + mux.HandleFunc("/api/v1/tunnel/batch-delete", h.tunnelBatchDelete) + mux.HandleFunc("/api/v1/tunnel/batch-redeploy", h.tunnelBatchRedeploy) + mux.HandleFunc("/api/v1/tunnel/user/assign", h.userTunnelAssign) + mux.HandleFunc("/api/v1/tunnel/user/batch-assign", h.userTunnelBatchAssign) + mux.HandleFunc("/api/v1/tunnel/user/remove", h.userTunnelRemove) + mux.HandleFunc("/api/v1/tunnel/user/update", h.userTunnelUpdate) mux.HandleFunc("/api/v1/forward/list", h.forwardList) + mux.HandleFunc("/api/v1/forward/create", h.forwardCreate) + mux.HandleFunc("/api/v1/forward/update", h.forwardUpdate) + mux.HandleFunc("/api/v1/forward/delete", h.forwardDelete) + mux.HandleFunc("/api/v1/forward/force-delete", h.forwardForceDelete) + mux.HandleFunc("/api/v1/forward/pause", h.forwardPause) + mux.HandleFunc("/api/v1/forward/resume", h.forwardResume) + mux.HandleFunc("/api/v1/forward/diagnose", h.forwardDiagnose) + mux.HandleFunc("/api/v1/forward/update-order", h.forwardUpdateOrder) + mux.HandleFunc("/api/v1/forward/batch-delete", h.forwardBatchDelete) + mux.HandleFunc("/api/v1/forward/batch-pause", h.forwardBatchPause) + mux.HandleFunc("/api/v1/forward/batch-resume", h.forwardBatchResume) + mux.HandleFunc("/api/v1/forward/batch-redeploy", h.forwardBatchRedeploy) + mux.HandleFunc("/api/v1/forward/batch-change-tunnel", h.forwardBatchChangeTunnel) mux.HandleFunc("/api/v1/speed-limit/list", h.speedLimitList) + mux.HandleFunc("/api/v1/speed-limit/create", h.speedLimitCreate) + mux.HandleFunc("/api/v1/speed-limit/update", h.speedLimitUpdate) + mux.HandleFunc("/api/v1/speed-limit/delete", h.speedLimitDelete) + mux.HandleFunc("/api/v1/speed-limit/tunnels", h.tunnelList) mux.HandleFunc("/api/v1/tunnel/user/tunnel", h.userTunnelVisibleList) mux.HandleFunc("/api/v1/tunnel/user/list", h.userTunnelList) mux.HandleFunc("/api/v1/group/tunnel/list", h.tunnelGroupList) + mux.HandleFunc("/api/v1/group/tunnel/create", h.groupTunnelCreate) + mux.HandleFunc("/api/v1/group/tunnel/update", h.groupTunnelUpdate) + mux.HandleFunc("/api/v1/group/tunnel/delete", h.groupTunnelDelete) + mux.HandleFunc("/api/v1/group/tunnel/assign", h.groupTunnelAssign) mux.HandleFunc("/api/v1/group/user/list", h.userGroupList) + mux.HandleFunc("/api/v1/group/user/create", h.groupUserCreate) + mux.HandleFunc("/api/v1/group/user/update", h.groupUserUpdate) + mux.HandleFunc("/api/v1/group/user/delete", h.groupUserDelete) + mux.HandleFunc("/api/v1/group/user/assign", h.groupUserAssign) mux.HandleFunc("/api/v1/group/permission/list", h.groupPermissionList) + mux.HandleFunc("/api/v1/group/permission/assign", h.groupPermissionAssign) + mux.HandleFunc("/api/v1/group/permission/remove", h.groupPermissionRemove) + mux.HandleFunc("/api/v1/open_api/sub_store", h.openAPISubStore) mux.HandleFunc("/flow/test", h.flowTest) mux.HandleFunc("/flow/config", h.flowConfig) mux.HandleFunc("/flow/upload", h.flowUpload) + mux.HandleFunc("/error", h.errorPage) } func (h *Handler) login(w http.ResponseWriter, r *http.Request) { @@ -240,11 +295,26 @@ func (h *Handler) forwardList(w http.ResponseWriter, r *http.Request) { return } + userID, roleID, err := userRoleFromRequest(r) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + items, err := h.repo.ListForwards() if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } + if roleID != 0 { + filtered := make([]map[string]interface{}, 0, len(items)) + for _, item := range items { + if asInt64(item["userId"], 0) == userID { + filtered = append(filtered, item) + } + } + items = filtered + } response.WriteJSON(w, response.OK(items)) } @@ -262,6 +332,92 @@ func (h *Handler) speedLimitList(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.OK(items)) } +func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + if h == nil || h.repo == nil || h.repo.DB() == nil { + response.WriteJSON(w, response.Err(-2, "database unavailable")) + return + } + + username := strings.TrimSpace(r.URL.Query().Get("user")) + password := strings.TrimSpace(r.URL.Query().Get("pwd")) + tunnel := strings.TrimSpace(r.URL.Query().Get("tunnel")) + if tunnel == "" { + tunnel = "-1" + } + + if username == "" { + response.WriteJSON(w, response.ErrDefault("用户不能为空")) + return + } + if password == "" { + response.WriteJSON(w, response.ErrDefault("密码不能为空")) + return + } + + user, err := h.repo.GetUserByUsername(username) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if user == nil || user.Pwd != security.MD5(password) { + response.WriteJSON(w, response.ErrDefault("鉴权失败")) + return + } + + const giga = int64(1024 * 1024 * 1024) + headerValue := "" + + if tunnel == "-1" { + headerValue = buildSubscriptionHeader(user.OutFlow, user.InFlow, user.Flow*giga, user.ExpTime/1000) + } else { + tunnelID, parseErr := strconv.ParseInt(tunnel, 10, 64) + if parseErr != nil || tunnelID <= 0 { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + + var userID int64 + var inFlow int64 + var outFlow int64 + var flow int64 + var expTime int64 + err = h.repo.DB().QueryRow(`SELECT user_id, in_flow, out_flow, flow, exp_time FROM user_tunnel WHERE id = ? LIMIT 1`, tunnelID). + Scan(&userID, &inFlow, &outFlow, &flow, &expTime) + if err != nil { + if err == sql.ErrNoRows { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if userID != user.ID { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + + headerValue = buildSubscriptionHeader(outFlow, inFlow, flow*giga, expTime/1000) + } + + w.Header().Set("subscription-userinfo", headerValue) + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + _, _ = w.Write([]byte(headerValue)) +} + +func (h *Handler) errorPage(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/html; charset=UTF-8") + w.WriteHeader(http.StatusNotFound) + _, _ = w.Write([]byte("错误 404
404
你推开了后端的大门,却发现里面只有寂寞。
")) +} + +func buildSubscriptionHeader(upload, download, total, expire int64) string { + return fmt.Sprintf("upload=%d; download=%d; total=%d; expire=%d", download, upload, total, expire) +} + func (h *Handler) userTunnelVisibleList(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { response.WriteJSON(w, response.ErrDefault("请求失败")) @@ -738,6 +894,18 @@ func userIDFromRequest(r *http.Request) (int64, error) { return parseUserID(claims.Sub) } +func userRoleFromRequest(r *http.Request) (int64, int, error) { + claims, ok := r.Context().Value(middleware.ClaimsContextKey).(auth.Claims) + if !ok { + return 0, 0, strconv.ErrSyntax + } + userID, err := parseUserID(claims.Sub) + if err != nil { + return 0, 0, err + } + return userID, claims.RoleID, nil +} + func nullableNullInt64(v sql.NullInt64) interface{} { if v.Valid { return v.Int64 diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go new file mode 100644 index 0000000..92dcfa1 --- /dev/null +++ b/go-backend/internal/http/handler/mutations.go @@ -0,0 +1,1990 @@ +package handler + +import ( + "crypto/rand" + "database/sql" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "net/http" + "sort" + "strconv" + "strings" + "time" + + "go-backend/internal/http/response" + "go-backend/internal/security" +) + +func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + + username := asString(req["user"]) + pwd := asString(req["pwd"]) + if username == "" || pwd == "" { + response.WriteJSON(w, response.ErrDefault("用户名或密码不能为空")) + return + } + + db := h.repo.DB() + if db == nil { + response.WriteJSON(w, response.Err(-2, "database unavailable")) + return + } + + var cnt int + if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ?`, username).Scan(&cnt); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if cnt > 0 { + response.WriteJSON(w, response.ErrDefault("用户名已存在")) + return + } + + status := asInt(req["status"], 1) + flow := asInt64(req["flow"], 100) + num := asInt(req["num"], 10) + expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()) + flowResetTime := asInt64(req["flowResetTime"], 1) + roleID := asInt(req["roleId"], asInt(req["role_id"], 1)) + now := time.Now().UnixMilli() + + _, err := db.Exec(` + INSERT INTO user(user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?, ?, ?) + `, username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, now, now, status) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + id := asInt64(req["id"], 0) + if id <= 0 { + response.WriteJSON(w, response.ErrDefault("用户ID不能为空")) + return + } + username := asString(req["user"]) + if username == "" { + response.WriteJSON(w, response.ErrDefault("用户名不能为空")) + return + } + + db := h.repo.DB() + if db == nil { + response.WriteJSON(w, response.Err(-2, "database unavailable")) + return + } + + var cnt int + if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ? AND id != ?`, username, id).Scan(&cnt); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if cnt > 0 { + response.WriteJSON(w, response.ErrDefault("用户名已存在")) + return + } + + flow := asInt64(req["flow"], 100) + num := asInt(req["num"], 10) + expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()) + flowResetTime := asInt64(req["flowResetTime"], 1) + status := asInt(req["status"], 1) + now := time.Now().UnixMilli() + + pwd := asString(req["pwd"]) + if strings.TrimSpace(pwd) == "" { + _, err := db.Exec(` + UPDATE user + SET user = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ?, updated_time = ? + WHERE id = ? + `, username, flow, num, expTime, flowResetTime, status, now, id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + } else { + _, err := db.Exec(` + UPDATE user + SET user = ?, pwd = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ?, updated_time = ? + WHERE id = ? + `, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, now, id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + } + + _, _ = db.Exec(`UPDATE user_tunnel SET flow = ?, num = ?, exp_time = ?, flow_reset_time = ? WHERE user_id = ?`, flow, num, expTime, flowResetTime, id) + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + id := idFromBody(r, w) + if id <= 0 { + return + } + + db := h.repo.DB() + tx, err := db.Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + + if _, err = tx.Exec(`DELETE FROM forward_port WHERE forward_id IN (SELECT id FROM forward WHERE user_id = ?)`, id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if _, err = tx.Exec(`DELETE FROM forward WHERE user_id = ?`, id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if _, err = tx.Exec(`DELETE FROM user_tunnel WHERE user_id = ?`, id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if _, err = tx.Exec(`DELETE FROM user_group_user WHERE user_id = ?`, id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if _, err = tx.Exec(`DELETE FROM user WHERE id = ?`, id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + if err = tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) userResetFlow(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + id := asInt64(req["id"], 0) + typeVal := asInt(req["type"], 0) + if id <= 0 || (typeVal != 1 && typeVal != 2) { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + + db := h.repo.DB() + if typeVal == 1 { + _, _ = db.Exec(`UPDATE user SET in_flow = 0, out_flow = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id) + _, _ = db.Exec(`UPDATE user_tunnel SET in_flow = 0, out_flow = 0 WHERE user_id = ?`, id) + } else { + _, _ = db.Exec(`UPDATE user_tunnel SET in_flow = 0, out_flow = 0 WHERE id = ?`, id) + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) captchaGenerate(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`)) + return + } + token := randomToken(16) + payload := map[string]interface{}{ + "id": token, + "data": map[string]interface{}{ + "id": token, + }, + "success": true, + } + w.Header().Set("Content-Type", "application/json; charset=utf-8") + _ = json.NewEncoder(w).Encode(payload) +} + +func (h *Handler) captchaVerify(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`)) + return + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`)) + return + } + id := asString(req["captchaId"]) + if id == "" { + id = asString(req["id"]) + } + payload := map[string]interface{}{ + "success": true, + "data": map[string]interface{}{"validToken": id}, + } + w.Header().Set("Content-Type", "application/json; charset=utf-8") + _ = json.NewEncoder(w).Encode(payload) +} + +func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + name := asString(req["name"]) + serverIP := asString(req["serverIp"]) + if name == "" || serverIP == "" { + response.WriteJSON(w, response.ErrDefault("节点名称和地址不能为空")) + return + } + + db := h.repo.DB() + now := time.Now().UnixMilli() + inx := nextIndex(db, "node") + _, err := db.Exec(` + INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, + name, + randomToken(16), + serverIP, + nullableText(asString(req["serverIpV4"])), + nullableText(asString(req["serverIpV6"])), + defaultString(asString(req["port"]), "1000-65535"), + nullableText(asString(req["interfaceName"])), + nullableText(""), + asInt(req["http"], 0), + asInt(req["tls"], 0), + asInt(req["socks"], 0), + now, + now, + 0, + defaultString(asString(req["tcpListenAddr"]), "[::]"), + defaultString(asString(req["udpListenAddr"]), "[::]"), + inx, + ) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + id := asInt64(req["id"], 0) + if id <= 0 { + response.WriteJSON(w, response.ErrDefault("节点ID不能为空")) + return + } + + var currentStatus int + var currentHTTP int + var currentTLS int + var currentSocks int + if err := h.repo.DB().QueryRow(`SELECT status, http, tls, socks FROM node WHERE id = ?`, id).Scan(¤tStatus, ¤tHTTP, ¤tTLS, ¤tSocks); err != nil { + if err == sql.ErrNoRows { + response.WriteJSON(w, response.ErrDefault("节点不存在")) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + newHTTP := asInt(req["http"], currentHTTP) + newTLS := asInt(req["tls"], currentTLS) + newSocks := asInt(req["socks"], currentSocks) + if currentStatus == 1 && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) { + if err := h.applyNodeProtocolChange(id, newHTTP, newTLS, newSocks); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + } + + now := time.Now().UnixMilli() + _, err := h.repo.DB().Exec(` + UPDATE node + SET name = ?, server_ip = ?, server_ip_v4 = ?, server_ip_v6 = ?, port = ?, interface_name = ?, http = ?, tls = ?, socks = ?, tcp_listen_addr = ?, udp_listen_addr = ?, updated_time = ? + WHERE id = ? + `, + asString(req["name"]), + asString(req["serverIp"]), + nullableText(asString(req["serverIpV4"])), + nullableText(asString(req["serverIpV6"])), + defaultString(asString(req["port"]), "1000-65535"), + nullableText(asString(req["interfaceName"])), + newHTTP, + newTLS, + newSocks, + defaultString(asString(req["tcpListenAddr"]), "[::]"), + defaultString(asString(req["udpListenAddr"]), "[::]"), + now, + id, + ) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) nodeDelete(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + id := idFromBody(r, w) + if id <= 0 { + return + } + if err := h.deleteNodeByID(id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) nodeInstall(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + id := idFromBody(r, w) + if id <= 0 { + return + } + db := h.repo.DB() + var secret string + if err := db.QueryRow(`SELECT secret FROM node WHERE id = ?`, id).Scan(&secret); err != nil { + response.WriteJSON(w, response.ErrDefault("节点不存在")) + return + } + var panelAddr string + if err := db.QueryRow(`SELECT value FROM vite_config WHERE name = 'ip' LIMIT 1`).Scan(&panelAddr); err != nil { + if err == sql.ErrNoRows { + response.WriteJSON(w, response.ErrDefault("请先前往网站配置中设置ip")) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + cmd := fmt.Sprintf("curl -L https://github.com/Sagit-chu/flux-panel/releases/latest/download/install.sh -o ./install.sh && chmod +x ./install.sh && ./install.sh -a %s -s %s", processServerAddress(panelAddr), secret) + response.WriteJSON(w, response.OK(cmd)) +} + +func (h *Handler) nodeUpdateOrder(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req struct { + Nodes []struct { + ID int64 `json:"id"` + Inx int `json:"inx"` + } `json:"nodes"` + } + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + for _, n := range req.Nodes { + _, _ = h.repo.DB().Exec(`UPDATE node SET inx = ?, updated_time = ? WHERE id = ?`, n.Inx, time.Now().UnixMilli(), n.ID) + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) nodeBatchDelete(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + ids := idsFromBody(r, w) + if ids == nil { + return + } + for _, id := range ids { + _ = h.deleteNodeByID(id) + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) nodeCheckStatus(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + items, err := h.repo.ListNodes() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OK(items)) +} + +func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + name := asString(req["name"]) + if name == "" { + response.WriteJSON(w, response.ErrDefault("隧道名称不能为空")) + return + } + typeVal := asInt(req["type"], 1) + flow := asInt64(req["flow"], 1) + status := asInt(req["status"], 1) + trafficRatio := asFloat(req["trafficRatio"], 1.0) + inIP := asString(req["inIp"]) + now := time.Now().UnixMilli() + inx := nextIndex(h.repo.DB(), "tunnel") + + tx, err := h.repo.DB().Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + + res, err := tx.Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + tunnelID, _ := res.LastInsertId() + if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) tunnelGet(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + id := idFromBody(r, w) + if id <= 0 { + return + } + items, err := h.repo.ListTunnels() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + for _, it := range items { + if asInt64(it["id"], 0) == id { + response.WriteJSON(w, response.OK(it)) + return + } + } + response.WriteJSON(w, response.ErrDefault("隧道不存在")) +} + +func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + id := asInt64(req["id"], 0) + if id <= 0 { + response.WriteJSON(w, response.ErrDefault("隧道ID不能为空")) + return + } + now := time.Now().UnixMilli() + _, err := h.repo.DB().Exec(`UPDATE tunnel SET name=?, type=?, flow=?, traffic_ratio=?, status=?, in_ip=?, updated_time=? WHERE id=?`, + asString(req["name"]), asInt(req["type"], 1), asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(asString(req["inIp"])), now, id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + tx, err := h.repo.DB().Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + if _, err := tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := replaceTunnelChainsTx(tx, id, req); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) tunnelDelete(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + id := idFromBody(r, w) + if id <= 0 { + return + } + if err := h.deleteTunnelByID(id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) tunnelDiagnose(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + id := asInt64FromBodyKey(r, w, "tunnelId") + if id <= 0 { + return + } + result, err := h.diagnoseTunnelRuntime(id) + if err != nil { + if strings.Contains(err.Error(), "不存在") || strings.Contains(err.Error(), "不完整") { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OK(result)) +} + +func (h *Handler) tunnelUpdateOrder(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + var req struct { + Tunnels []struct { + ID int64 `json:"id"` + Inx int `json:"inx"` + } `json:"tunnels"` + } + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + for _, t := range req.Tunnels { + _, _ = h.repo.DB().Exec(`UPDATE tunnel SET inx = ?, updated_time = ? WHERE id = ?`, t.Inx, time.Now().UnixMilli(), t.ID) + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) { + ids := idsFromBody(r, w) + if ids == nil { + return + } + success := 0 + fail := 0 + for _, id := range ids { + if err := h.deleteTunnelByID(id); err != nil { + fail++ + } else { + success++ + } + } + response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail})) +} + +func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) { + ids := idsFromBody(r, w) + if ids == nil { + return + } + success := 0 + fail := 0 + for _, tunnelID := range ids { + forwards, err := h.listForwardsByTunnel(tunnelID) + if err != nil { + fail++ + continue + } + if len(forwards) == 0 { + success++ + continue + } + ok := true + for i := range forwards { + if err := h.syncForwardServices(&forwards[i], "UpdateService", true); err != nil { + ok = false + break + } + } + if ok { + success++ + } else { + fail++ + } + } + response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail})) +} + +func (h *Handler) userTunnelAssign(w http.ResponseWriter, r *http.Request) { + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + if err := h.upsertUserTunnel(req); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) userTunnelBatchAssign(w http.ResponseWriter, r *http.Request) { + var req struct { + UserID int64 `json:"userId"` + Tunnels []struct { + TunnelID int64 `json:"tunnelId"` + SpeedID *int64 `json:"speedId"` + } `json:"tunnels"` + } + if err := decodeJSON(r.Body, &req); err != nil || req.UserID <= 0 { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + for _, t := range req.Tunnels { + m := map[string]interface{}{"userId": req.UserID, "tunnelId": t.TunnelID} + if t.SpeedID != nil { + m["speedId"] = *t.SpeedID + } + if err := h.upsertUserTunnel(m); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) userTunnelRemove(w http.ResponseWriter, r *http.Request) { + id := idFromBody(r, w) + if id <= 0 { + return + } + _, err := h.repo.DB().Exec(`DELETE FROM user_tunnel WHERE id = ?`, id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) { + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + id := asInt64(req["id"], 0) + if id <= 0 { + response.WriteJSON(w, response.ErrDefault("权限ID不能为空")) + return + } + _, err := h.repo.DB().Exec(` + UPDATE user_tunnel SET flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, speed_id = ?, status = ? WHERE id = ? + `, + asInt64(req["flow"], 0), + asInt(req["num"], 0), + asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()), + asInt64(req["flowResetTime"], 1), + nullableInt(asAnyToInt64Ptr(req["speedId"])), + asInt(req["status"], 1), + id, + ) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) { + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + userID, roleID, err := userRoleFromRequest(r) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + tunnelID := asInt64(req["tunnelId"], 0) + if tunnelID <= 0 { + response.WriteJSON(w, response.ErrDefault("隧道ID不能为空")) + return + } + if err := h.ensureTunnelPermission(userID, roleID, tunnelID); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + tunnel, err := h.getTunnelRecord(tunnelID) + if err != nil { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + if tunnel.Status != 1 { + response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法创建转发")) + return + } + name := asString(req["name"]) + remoteAddr := asString(req["remoteAddr"]) + if name == "" || remoteAddr == "" { + response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空")) + return + } + port := asInt(req["inPort"], 0) + if port <= 0 { + port = h.pickTunnelPort(tunnelID) + } + if port <= 0 { + port = 10000 + } + now := time.Now().UnixMilli() + inx := nextIndex(h.repo.DB(), "forward") + var userName string + _ = h.repo.DB().QueryRow(`SELECT user FROM user WHERE id = ?`, userID).Scan(&userName) + if userName == "" { + userName = "user" + } + tx, err := h.repo.DB().Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + res, err := tx.Exec(` + INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?) + `, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + forwardID, _ := res.LastInsertId() + entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID) + for _, nodeID := range entryNodes { + _, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port) + } + if err := tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + createdForward, err := h.getForwardRecord(forwardID) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := h.syncForwardServices(createdForward, "AddService", false); err != nil { + _ = h.deleteForwardByID(forwardID) + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + id := asInt64(req["id"], 0) + if id <= 0 { + response.WriteJSON(w, response.ErrDefault("转发ID不能为空")) + return + } + forward, actorUserID, actorRole, err := h.resolveForwardAccess(r, id) + if err != nil { + if errors.Is(err, errForwardNotFound) { + response.WriteJSON(w, response.ErrDefault("转发不存在")) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + tunnelID := asInt64(req["tunnelId"], forward.TunnelID) + if tunnelID <= 0 { + response.WriteJSON(w, response.ErrDefault("隧道ID不能为空")) + return + } + if err := h.ensureTunnelPermission(actorUserID, actorRole, tunnelID); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + tunnel, err := h.getTunnelRecord(tunnelID) + if err != nil { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + if tunnel.Status != 1 { + response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法更新转发")) + return + } + + name := strings.TrimSpace(asString(req["name"])) + if name == "" { + name = forward.Name + } + remoteAddr := strings.TrimSpace(asString(req["remoteAddr"])) + if remoteAddr == "" { + remoteAddr = forward.RemoteAddr + } + strategy := strings.TrimSpace(asString(req["strategy"])) + if strategy == "" { + strategy = forward.Strategy + } + + port := asInt(req["inPort"], 0) + if port <= 0 { + var minPort sql.NullInt64 + _ = h.repo.DB().QueryRow(`SELECT MIN(port) FROM forward_port WHERE forward_id = ?`, id).Scan(&minPort) + if minPort.Valid { + port = int(minPort.Int64) + } + if port <= 0 { + port = h.pickTunnelPort(tunnelID) + } + } + now := time.Now().UnixMilli() + _, err = h.repo.DB().Exec(` + UPDATE forward SET name = ?, tunnel_id = ?, remote_addr = ?, strategy = ?, updated_time = ? WHERE id = ? + `, name, tunnelID, remoteAddr, strategy, now, id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + _ = h.replaceForwardPorts(id, tunnelID, port) + updatedForward, err := h.getForwardRecord(id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) { + id := idFromBody(r, w) + if id <= 0 { + return + } + forward, _, _, err := h.resolveForwardAccess(r, id) + if err != nil { + if errors.Is(err, errForwardNotFound) { + response.WriteJSON(w, response.ErrDefault("转发不存在")) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := h.controlForwardServices(forward, "DeleteService", true); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + if err := h.deleteForwardByID(id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) { + h.forwardDelete(w, r) +} + +func (h *Handler) forwardPause(w http.ResponseWriter, r *http.Request) { + id := idFromBody(r, w) + if id <= 0 { + return + } + forward, _, _, err := h.resolveForwardAccess(r, id) + if err != nil { + if errors.Is(err, errForwardNotFound) { + response.WriteJSON(w, response.ErrDefault("转发不存在")) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := h.controlForwardServices(forward, "PauseService", false); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + _, _ = h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id) + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) forwardResume(w http.ResponseWriter, r *http.Request) { + id := idFromBody(r, w) + if id <= 0 { + return + } + forward, _, _, err := h.resolveForwardAccess(r, id) + if err != nil { + if errors.Is(err, errForwardNotFound) { + response.WriteJSON(w, response.ErrDefault("转发不存在")) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := h.controlForwardServices(forward, "ResumeService", false); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + _, _ = h.repo.DB().Exec(`UPDATE forward SET status = 1, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id) + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) forwardDiagnose(w http.ResponseWriter, r *http.Request) { + id := asInt64FromBodyKey(r, w, "forwardId") + if id <= 0 { + return + } + forward, _, _, err := h.resolveForwardAccess(r, id) + if err != nil { + if errors.Is(err, errForwardNotFound) { + response.WriteJSON(w, response.ErrDefault("转发不存在")) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + payload, err := h.diagnoseForwardRuntime(forward) + if err != nil { + if strings.Contains(err.Error(), "不存在") || strings.Contains(err.Error(), "不能为空") || strings.Contains(err.Error(), "错误") { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OK(payload)) +} + +func (h *Handler) forwardUpdateOrder(w http.ResponseWriter, r *http.Request) { + var req struct { + Forwards []struct { + ID int64 `json:"id"` + Inx int `json:"inx"` + } `json:"forwards"` + } + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + for _, f := range req.Forwards { + _, _ = h.repo.DB().Exec(`UPDATE forward SET inx = ?, updated_time = ? WHERE id = ?`, f.Inx, time.Now().UnixMilli(), f.ID) + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) { + ids := idsFromBody(r, w) + if ids == nil { + return + } + actorUserID, actorRole, err := userRoleFromRequest(r) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + s := 0 + f := 0 + for _, id := range ids { + forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id) + if accessErr != nil { + f++ + continue + } + if err := h.controlForwardServices(forward, "DeleteService", true); err != nil { + f++ + continue + } + if err := h.deleteForwardByID(id); err != nil { + f++ + } else { + s++ + } + } + response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f})) +} + +func (h *Handler) forwardBatchPause(w http.ResponseWriter, r *http.Request) { + ids := idsFromBody(r, w) + if ids == nil { + return + } + actorUserID, actorRole, err := userRoleFromRequest(r) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + s := 0 + f := 0 + for _, id := range ids { + forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id) + if accessErr != nil { + f++ + continue + } + if err := h.controlForwardServices(forward, "PauseService", false); err != nil { + f++ + continue + } + if _, err := h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id); err != nil { + f++ + } else { + s++ + } + } + response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f})) +} + +func (h *Handler) forwardBatchResume(w http.ResponseWriter, r *http.Request) { + ids := idsFromBody(r, w) + if ids == nil { + return + } + actorUserID, actorRole, err := userRoleFromRequest(r) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + s := 0 + f := 0 + for _, id := range ids { + forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id) + if accessErr != nil { + f++ + continue + } + if err := h.controlForwardServices(forward, "ResumeService", false); err != nil { + f++ + continue + } + if _, err := h.repo.DB().Exec(`UPDATE forward SET status = 1, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id); err != nil { + f++ + } else { + s++ + } + } + response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f})) +} + +func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) { + ids := idsFromBody(r, w) + if ids == nil { + return + } + actorUserID, actorRole, err := userRoleFromRequest(r) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + s := 0 + f := 0 + for _, id := range ids { + forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id) + if accessErr != nil { + f++ + continue + } + if err := h.syncForwardServices(forward, "UpdateService", true); err != nil { + f++ + } else { + s++ + } + } + response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f})) +} + +func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Request) { + var req struct { + ForwardIDs []int64 `json:"forwardIds"` + TargetTunnelID int64 `json:"targetTunnelId"` + } + if err := decodeJSON(r.Body, &req); err != nil || req.TargetTunnelID <= 0 { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + actorUserID, actorRole, err := userRoleFromRequest(r) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + if err := h.ensureTunnelPermission(actorUserID, actorRole, req.TargetTunnelID); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + targetTunnel, err := h.getTunnelRecord(req.TargetTunnelID) + if err != nil { + response.WriteJSON(w, response.ErrDefault("目标隧道不存在")) + return + } + if targetTunnel.Status != 1 { + response.WriteJSON(w, response.ErrDefault("目标隧道已禁用")) + return + } + success := 0 + fail := 0 + for _, id := range req.ForwardIDs { + if id <= 0 { + continue + } + forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id) + if accessErr != nil { + fail++ + continue + } + if forward.TunnelID == req.TargetTunnelID { + fail++ + continue + } + var port sql.NullInt64 + _ = h.repo.DB().QueryRow(`SELECT MIN(port) FROM forward_port WHERE forward_id = ?`, id).Scan(&port) + _, err := h.repo.DB().Exec(`UPDATE forward SET tunnel_id = ?, updated_time = ? WHERE id = ?`, req.TargetTunnelID, time.Now().UnixMilli(), id) + if err != nil { + fail++ + continue + } + p := 0 + if port.Valid { + p = int(port.Int64) + } + if p <= 0 { + p = h.pickTunnelPort(req.TargetTunnelID) + } + _ = h.replaceForwardPorts(id, req.TargetTunnelID, p) + updatedForward, fetchErr := h.getForwardRecord(id) + if fetchErr != nil { + fail++ + continue + } + if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil { + fail++ + continue + } + success++ + } + response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail})) +} + +func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) { + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + tunnelID := asInt64(req["tunnelId"], 0) + if tunnelID <= 0 { + response.WriteJSON(w, response.ErrDefault("隧道ID不能为空")) + return + } + name := asString(req["name"]) + if name == "" { + response.WriteJSON(w, response.ErrDefault("名称不能为空")) + return + } + var tunnelName string + _ = h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName) + if tunnelName == "" { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + now := time.Now().UnixMilli() + _, err := h.repo.DB().Exec(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?)`, + name, asInt(req["speed"], 100), tunnelID, tunnelName, now, now, asInt(req["status"], 1)) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) { + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + id := asInt64(req["id"], 0) + tunnelID := asInt64(req["tunnelId"], 0) + if id <= 0 || tunnelID <= 0 { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + var tunnelName string + _ = h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName) + if tunnelName == "" { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + _, err := h.repo.DB().Exec(`UPDATE speed_limit SET name=?, speed=?, tunnel_id=?, tunnel_name=?, status=?, updated_time=? WHERE id=?`, + asString(req["name"]), asInt(req["speed"], 100), tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli(), id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) { + id := idFromBody(r, w) + if id <= 0 { + return + } + _, err := h.repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) groupTunnelCreate(w http.ResponseWriter, r *http.Request) { + h.groupCreate(w, r, "tunnel_group") +} + +func (h *Handler) groupTunnelUpdate(w http.ResponseWriter, r *http.Request) { + h.groupUpdate(w, r, "tunnel_group") +} + +func (h *Handler) groupTunnelDelete(w http.ResponseWriter, r *http.Request) { + h.groupDelete(w, r, "tunnel_group") +} + +func (h *Handler) groupUserCreate(w http.ResponseWriter, r *http.Request) { + h.groupCreate(w, r, "user_group") +} + +func (h *Handler) groupUserUpdate(w http.ResponseWriter, r *http.Request) { + h.groupUpdate(w, r, "user_group") +} + +func (h *Handler) groupUserDelete(w http.ResponseWriter, r *http.Request) { + h.groupDelete(w, r, "user_group") +} + +func (h *Handler) groupTunnelAssign(w http.ResponseWriter, r *http.Request) { + var req struct { + GroupID int64 `json:"groupId"` + TunnelIDs []int64 `json:"tunnelIds"` + } + if err := decodeJSON(r.Body, &req); err != nil || req.GroupID <= 0 { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + tx, err := h.repo.DB().Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + _, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, req.GroupID) + for _, tid := range req.TunnelIDs { + _, _ = tx.Exec(`INSERT OR IGNORE INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, req.GroupID, tid, time.Now().UnixMilli()) + } + if err := tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + _ = h.syncPermissionsByTunnelGroup(req.GroupID) + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) { + var req struct { + GroupID int64 `json:"groupId"` + UserIDs []int64 `json:"userIds"` + } + if err := decodeJSON(r.Body, &req); err != nil || req.GroupID <= 0 { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + tx, err := h.repo.DB().Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID) + for _, uid := range req.UserIDs { + _, _ = tx.Exec(`INSERT OR IGNORE INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?)`, req.GroupID, uid, time.Now().UnixMilli()) + } + if err := tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + _ = h.syncPermissionsByUserGroup(req.GroupID) + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) groupPermissionAssign(w http.ResponseWriter, r *http.Request) { + var req struct { + UserGroupID int64 `json:"userGroupId"` + TunnelGroupID int64 `json:"tunnelGroupId"` + } + if err := decodeJSON(r.Body, &req); err != nil || req.UserGroupID <= 0 || req.TunnelGroupID <= 0 { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + _, err := h.repo.DB().Exec(`INSERT OR IGNORE INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?)`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli()) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + _ = h.applyGroupPermission(req.UserGroupID, req.TunnelGroupID) + response.WriteJSON(w, response.OK("权限分配成功")) +} + +func (h *Handler) groupPermissionRemove(w http.ResponseWriter, r *http.Request) { + id := idFromBody(r, w) + if id <= 0 { + return + } + var ug, tg int64 + _ = h.repo.DB().QueryRow(`SELECT user_group_id, tunnel_group_id FROM group_permission WHERE id = ?`, id).Scan(&ug, &tg) + _, _ = h.repo.DB().Exec(`DELETE FROM group_permission WHERE id = ?`, id) + _, _ = h.repo.DB().Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, ug, tg) + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) groupCreate(w http.ResponseWriter, r *http.Request, table string) { + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + name := asString(req["name"]) + if name == "" { + response.WriteJSON(w, response.ErrDefault("分组名称不能为空")) + return + } + now := time.Now().UnixMilli() + _, err := h.repo.DB().Exec(`INSERT INTO `+table+`(name, created_time, updated_time, status) VALUES(?, ?, ?, ?)`, name, now, now, asInt(req["status"], 1)) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) groupUpdate(w http.ResponseWriter, r *http.Request, table string) { + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return + } + id := asInt64(req["id"], 0) + if id <= 0 { + response.WriteJSON(w, response.ErrDefault("分组ID不能为空")) + return + } + _, err := h.repo.DB().Exec(`UPDATE `+table+` SET name = ?, status = ?, updated_time = ? WHERE id = ?`, asString(req["name"]), asInt(req["status"], 1), time.Now().UnixMilli(), id) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) groupDelete(w http.ResponseWriter, r *http.Request, table string) { + id := idFromBody(r, w) + if id <= 0 { + return + } + tx, err := h.repo.DB().Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + if table == "tunnel_group" { + _, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, id) + _, _ = tx.Exec(`DELETE FROM group_permission WHERE tunnel_group_id = ?`, id) + _, _ = tx.Exec(`DELETE FROM group_permission_grant WHERE tunnel_group_id = ?`, id) + } else { + _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, id) + _, _ = tx.Exec(`DELETE FROM group_permission WHERE user_group_id = ?`, id) + _, _ = tx.Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ?`, id) + } + _, _ = tx.Exec(`DELETE FROM `+table+` WHERE id = ?`, id) + if err := tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error { + db := h.repo.DB() + userIDs, _ := queryInt64List(db, `SELECT user_id FROM user_group_user WHERE user_group_id = ?`, userGroupID) + tunnelIDs, _ := queryInt64List(db, `SELECT tunnel_id FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tunnelGroupID) + for _, uid := range userIDs { + for _, tid := range tunnelIDs { + utID, created, err := ensureUserTunnelGrant(db, uid, tid) + if err != nil { + continue + } + createdByGroup := 0 + if created { + createdByGroup = 1 + } + _, _ = db.Exec(`INSERT OR IGNORE INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?)`, + userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli()) + } + } + return nil +} + +func (h *Handler) syncPermissionsByUserGroup(userGroupID int64) error { + db := h.repo.DB() + pairs, err := queryPairs(db, `SELECT user_group_id, tunnel_group_id FROM group_permission WHERE user_group_id = ?`, userGroupID) + if err != nil { + return err + } + for _, p := range pairs { + _ = h.applyGroupPermission(p[0], p[1]) + } + return nil +} + +func (h *Handler) syncPermissionsByTunnelGroup(tunnelGroupID int64) error { + db := h.repo.DB() + pairs, err := queryPairs(db, `SELECT user_group_id, tunnel_group_id FROM group_permission WHERE tunnel_group_id = ?`, tunnelGroupID) + if err != nil { + return err + } + for _, p := range pairs { + _ = h.applyGroupPermission(p[0], p[1]) + } + return nil +} + +func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, error) { + var id int64 + err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&id) + if err == nil { + return id, false, nil + } + if err != sql.ErrNoRows { + return 0, false, err + } + var flow int64 + var num int + var expTime int64 + var flowReset int64 + if err := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset); err != nil { + return 0, false, err + } + res, err := db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1)`, + userID, tunnelID, num, flow, flowReset, expTime) + if err != nil { + return 0, false, err + } + id, _ = res.LastInsertId() + return id, true, nil +} + +func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error) { + rows, err := db.Query(q, args...) + if err != nil { + return nil, err + } + defer rows.Close() + out := make([]int64, 0) + for rows.Next() { + var v int64 + if err := rows.Scan(&v); err != nil { + return nil, err + } + out = append(out, v) + } + return out, rows.Err() +} + +func queryPairs(db *sql.DB, q string, args ...interface{}) ([][2]int64, error) { + rows, err := db.Query(q, args...) + if err != nil { + return nil, err + } + defer rows.Close() + out := make([][2]int64, 0) + for rows.Next() { + var a, b int64 + if err := rows.Scan(&a, &b); err != nil { + return nil, err + } + out = append(out, [2]int64{a, b}) + } + return out, rows.Err() +} + +func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{}) error { + inNodes := asMapSlice(req["inNodeId"]) + for _, n := range inNodes { + nodeID := asInt64(n["nodeId"], 0) + if nodeID <= 0 { + continue + } + _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, NULL, NULL, 0, ?)`, + tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls")) + if err != nil { + return err + } + } + for _, n := range asMapSlice(req["outNodeId"]) { + nodeID := asInt64(n["nodeId"], 0) + if nodeID <= 0 { + continue + } + _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, NULL, NULL, 0, ?)`, + tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls")) + if err != nil { + return err + } + } + chainNodes := asAnySlice(req["chainNodes"]) + for i, grp := range chainNodes { + for _, n := range asMapSlice(grp) { + nodeID := asInt64(n["nodeId"], 0) + if nodeID <= 0 { + continue + } + _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, NULL, ?, ?, ?)`, + tunnelID, nodeID, defaultString(asString(n["strategy"]), "round"), i, defaultString(asString(n["protocol"]), "tls")) + if err != nil { + return err + } + } + } + return nil +} + +func (h *Handler) deleteNodeByID(id int64) error { + tx, err := h.repo.DB().Begin() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + _, _ = tx.Exec(`DELETE FROM forward_port WHERE node_id = ?`, id) + _, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE node_id = ?`, id) + _, err = tx.Exec(`DELETE FROM node WHERE id = ?`, id) + if err != nil { + return err + } + return tx.Commit() +} + +func (h *Handler) deleteTunnelByID(id int64) error { + tx, err := h.repo.DB().Begin() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + _, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id IN (SELECT id FROM forward WHERE tunnel_id = ?)`, id) + _, _ = tx.Exec(`DELETE FROM forward WHERE tunnel_id = ?`, id) + _, _ = tx.Exec(`DELETE FROM user_tunnel WHERE tunnel_id = ?`, id) + _, _ = tx.Exec(`DELETE FROM speed_limit WHERE tunnel_id = ?`, id) + _, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id) + _, err = tx.Exec(`DELETE FROM tunnel WHERE id = ?`, id) + if err != nil { + return err + } + return tx.Commit() +} + +func (h *Handler) deleteForwardByID(id int64) error { + tx, err := h.repo.DB().Begin() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + _, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, id) + _, err = tx.Exec(`DELETE FROM forward WHERE id = ?`, id) + if err != nil { + return err + } + return tx.Commit() +} + +func (h *Handler) batchForwardDelete(ids []int64) (int, int) { + s := 0 + f := 0 + for _, id := range ids { + if err := h.deleteForwardByID(id); err != nil { + f++ + } else { + s++ + } + } + return s, f +} + +func (h *Handler) batchForwardStatus(ids []int64, status int) (int, int) { + s := 0 + f := 0 + for _, id := range ids { + if _, err := h.repo.DB().Exec(`UPDATE forward SET status = ?, updated_time = ? WHERE id = ?`, status, time.Now().UnixMilli(), id); err != nil { + f++ + } else { + s++ + } + } + return s, f +} + +func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) { + rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 ORDER BY inx ASC, id ASC`, tunnelID) + if err != nil { + return nil, err + } + defer rows.Close() + out := make([]int64, 0) + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err == nil { + out = append(out, id) + } + } + return out, rows.Err() +} + +func (h *Handler) pickTunnelPort(tunnelID int64) int { + entry, _ := h.tunnelEntryNodeIDs(tunnelID) + if len(entry) == 0 { + return 10000 + } + var portRange string + _ = h.repo.DB().QueryRow(`SELECT port FROM node WHERE id = ?`, entry[0]).Scan(&portRange) + if portRange == "" { + return 10000 + } + first := strings.Split(portRange, ",")[0] + first = strings.TrimSpace(first) + if strings.Contains(first, "-") { + parts := strings.SplitN(first, "-", 2) + p, _ := strconv.Atoi(strings.TrimSpace(parts[0])) + if p > 0 { + return p + } + } + if p, err := strconv.Atoi(first); err == nil && p > 0 { + return p + } + return 10000 +} + +func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int) error { + tx, err := h.repo.DB().Begin() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + _, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID) + entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID) + for _, nodeID := range entryNodes { + _, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port) + } + return tx.Commit() +} + +func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { + userID := asInt64(req["userId"], 0) + tunnelID := asInt64(req["tunnelId"], 0) + if userID <= 0 || tunnelID <= 0 { + return fmt.Errorf("userId or tunnelId missing") + } + db := h.repo.DB() + var existingID int64 + err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&existingID) + flow := asInt64(req["flow"], -1) + num := asInt(req["num"], -1) + expTime := asInt64(req["expTime"], -1) + flowReset := asInt64(req["flowResetTime"], -1) + status := asInt(req["status"], 1) + speedID := asAnyToInt64Ptr(req["speedId"]) + if err == sql.ErrNoRows { + if flow < 0 || num < 0 || expTime < 0 || flowReset < 0 { + _ = db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset) + } + _, err = db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?)`, + userID, tunnelID, nullableInt(speedID), num, flow, flowReset, expTime, status) + return err + } + if err != nil { + return err + } + if flow < 0 { + flow = 0 + } + if num < 0 { + num = 0 + } + if expTime < 0 { + expTime = time.Now().Add(365 * 24 * time.Hour).UnixMilli() + } + if flowReset < 0 { + flowReset = 1 + } + _, err = db.Exec(`UPDATE user_tunnel SET speed_id = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ? WHERE id = ?`, + nullableInt(speedID), flow, num, expTime, flowReset, status, existingID) + return err +} + +func asAnySlice(v interface{}) []interface{} { + if v == nil { + return nil + } + if arr, ok := v.([]interface{}); ok { + return arr + } + return nil +} + +func asMapSlice(v interface{}) []map[string]interface{} { + arr := asAnySlice(v) + if arr == nil { + return nil + } + out := make([]map[string]interface{}, 0, len(arr)) + for _, it := range arr { + if m, ok := it.(map[string]interface{}); ok { + out = append(out, m) + } + } + return out +} + +func asString(v interface{}) string { + switch t := v.(type) { + case nil: + return "" + case string: + return strings.TrimSpace(t) + case float64: + if t == float64(int64(t)) { + return strconv.FormatInt(int64(t), 10) + } + return strconv.FormatFloat(t, 'f', -1, 64) + case int, int32, int64: + return fmt.Sprintf("%v", t) + default: + b, _ := json.Marshal(t) + return strings.Trim(string(b), "\"") + } +} + +func asInt(v interface{}, def int) int { + s := asString(v) + if s == "" { + return def + } + i, err := strconv.Atoi(s) + if err != nil { + return def + } + return i +} + +func asInt64(v interface{}, def int64) int64 { + s := asString(v) + if s == "" { + return def + } + i, err := strconv.ParseInt(s, 10, 64) + if err != nil { + return def + } + return i +} + +func asFloat(v interface{}, def float64) float64 { + s := asString(v) + if s == "" { + return def + } + f, err := strconv.ParseFloat(s, 64) + if err != nil { + return def + } + return f +} + +func asAnyToInt64Ptr(v interface{}) *int64 { + s := asString(v) + if s == "" || strings.EqualFold(s, "null") { + return nil + } + i, err := strconv.ParseInt(s, 10, 64) + if err != nil { + return nil + } + return &i +} + +func idFromBody(r *http.Request, w http.ResponseWriter) int64 { + return asInt64FromBodyKey(r, w, "id") +} + +func asInt64FromBodyKey(r *http.Request, w http.ResponseWriter, key string) int64 { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return 0 + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return 0 + } + id := asInt64(req[key], 0) + if id <= 0 { + response.WriteJSON(w, response.ErrDefault("参数错误")) + return 0 + } + return id +} + +func idsFromBody(r *http.Request, w http.ResponseWriter) []int64 { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return nil + } + var req map[string]interface{} + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("请求参数错误")) + return nil + } + arr := asAnySlice(req["ids"]) + if len(arr) == 0 { + response.WriteJSON(w, response.ErrDefault("ids不能为空")) + return nil + } + ids := make([]int64, 0, len(arr)) + for _, x := range arr { + id := asInt64(x, 0) + if id > 0 { + ids = append(ids, id) + } + } + sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) + return ids +} + +func nullableText(s string) interface{} { + if strings.TrimSpace(s) == "" { + return nil + } + return s +} + +func nullableInt(v *int64) interface{} { + if v == nil { + return nil + } + return *v +} + +func defaultString(v, def string) string { + if strings.TrimSpace(v) == "" { + return def + } + return v +} + +func randomToken(n int) string { + buf := make([]byte, n) + if _, err := rand.Read(buf); err != nil { + return strconv.FormatInt(time.Now().UnixNano(), 16) + } + return hex.EncodeToString(buf) +} + +func nextIndex(db *sql.DB, table string) int { + if db == nil { + return 0 + } + row := db.QueryRow(`SELECT COALESCE(MAX(inx), -1) + 1 FROM ` + table) + var n int + if err := row.Scan(&n); err != nil { + return 0 + } + if n < 0 { + return 0 + } + return n +} diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index a4fa678..8fd5806 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -25,6 +25,13 @@ type Repository struct { db *sql.DB } +func (r *Repository) DB() *sql.DB { + if r == nil { + return nil + } + return r.db +} + type User struct { ID int64 User string diff --git a/go-backend/internal/ws/server.go b/go-backend/internal/ws/server.go index a12d3db..f36ec97 100644 --- a/go-backend/internal/ws/server.go +++ b/go-backend/internal/ws/server.go @@ -2,11 +2,14 @@ package ws import ( "encoding/json" + "errors" + "fmt" "log" "net/http" "strconv" "strings" "sync" + "time" "github.com/gorilla/websocket" @@ -38,15 +41,36 @@ type nodeSession struct { conn *connWrap } +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 +} + +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 *sqlite.Repository jwtSecret string upgrader websocket.Upgrader - mu sync.RWMutex - admins map[*connWrap]struct{} - nodes map[int64]*nodeSession - byConn map[*websocket.Conn]*nodeSession + mu sync.RWMutex + admins map[*connWrap]struct{} + nodes map[int64]*nodeSession + byConn map[*websocket.Conn]*nodeSession + pending map[string]pendingRequest } func NewServer(repo *sqlite.Repository, jwtSecret string) *Server { @@ -56,9 +80,10 @@ func NewServer(repo *sqlite.Repository, jwtSecret string) *Server { 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), + admins: make(map[*connWrap]struct{}), + nodes: make(map[int64]*nodeSession), + byConn: make(map[*websocket.Conn]*nodeSession), + pending: make(map[string]pendingRequest), } } @@ -150,6 +175,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64 delete(s.byConn, conn) s.mu.Unlock() if needOfflineBroadcast { + s.failPendingForNode(nodeID, "节点连接已断开") _ = s.repo.UpdateNodeStatus(nodeID, 0) s.broadcastStatus(nodeID, 0) } @@ -163,10 +189,186 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64 } msg := decryptIfNeeded(payload, secret) + s.tryResolvePending(nodeID, msg) s.broadcastInfo(nodeID, msg) } } +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 strings.TrimSpace(ns.secret) != "" { + crypto, err := security.NewAESCrypto(ns.secret) + if err != nil { + cleanup() + return CommandResult{}, err + } + encrypted, err := 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() + err = ns.conn.conn.WriteMessage(websocket.TextMessage, messageData) + 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 + } + + 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), diff --git a/go-backend/tests/contract/forward_contract_test.go b/go-backend/tests/contract/forward_contract_test.go new file mode 100644 index 0000000..0b9973f --- /dev/null +++ b/go-backend/tests/contract/forward_contract_test.go @@ -0,0 +1,128 @@ +package contract_test + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "testing" + "time" + + "go-backend/internal/auth" + "go-backend/internal/http/response" +) + +func TestForwardOwnershipAndScopeContracts(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + if _, err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now); err != nil { + t.Fatalf("insert user: %v", err) + } + + res, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "contract-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, err := res.LastInsertId() + if err != nil { + t.Fatalf("get tunnel id: %v", err) + } + + resAdmin, err := repo.DB().Exec(` + INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?) + `, 1, "admin_user", "admin-forward", tunnelID, "1.1.1.1:443", "fifo", now, now, 0) + if err != nil { + t.Fatalf("insert admin forward: %v", err) + } + adminForwardID, err := resAdmin.LastInsertId() + if err != nil { + t.Fatalf("get admin forward id: %v", err) + } + + resUser, err := repo.DB().Exec(` + INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?) + `, 2, "normal_user", "user-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 1) + if err != nil { + t.Fatalf("insert user forward: %v", err) + } + userForwardID, err := resUser.LastInsertId() + if err != nil { + t.Fatalf("get user forward id: %v", err) + } + + userToken, err := auth.GenerateToken(2, "normal_user", 1, secret) + if err != nil { + t.Fatalf("generate user token: %v", err) + } + + t.Run("non-owner cannot delete another user's forward", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/delete", bytes.NewBufferString(`{"id":`+jsonNumber(adminForwardID)+`}`)) + req.Header.Set("Authorization", userToken) + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + assertCodeMsg(t, res, -1, "转发不存在") + }) + + t.Run("non-admin forward list is scoped to owner", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/list", bytes.NewBufferString(`{}`)) + req.Header.Set("Authorization", userToken) + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code != 0 { + t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg) + } + arr, ok := out.Data.([]interface{}) + if !ok { + t.Fatalf("expected array data, got %T", out.Data) + } + if len(arr) != 1 { + t.Fatalf("expected 1 forward, got %d", len(arr)) + } + item, ok := arr[0].(map[string]interface{}) + if !ok { + t.Fatalf("expected object item, got %T", arr[0]) + } + if got := int64(item["id"].(float64)); got != userForwardID { + t.Fatalf("expected forward id %d, got %d", userForwardID, got) + } + }) + + t.Run("diagnose no longer returns hardcoded success", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/diagnose", bytes.NewBufferString(`{"forwardId":`+jsonNumber(userForwardID)+`}`)) + req.Header.Set("Authorization", userToken) + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code == 0 { + t.Fatalf("expected non-zero code for missing runtime path, got success") + } + }) +} + +func jsonNumber(v int64) string { + return strconv.FormatInt(v, 10) +} diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go new file mode 100644 index 0000000..becb59e --- /dev/null +++ b/go-backend/tests/contract/migration_contract_test.go @@ -0,0 +1,154 @@ +package contract_test + +import ( + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "path/filepath" + "strconv" + "strings" + "testing" + "time" + + "go-backend/internal/auth" + httpserver "go-backend/internal/http" + "go-backend/internal/http/handler" + "go-backend/internal/http/response" + "go-backend/internal/store/sqlite" +) + +func TestOpenAPISubStoreContracts(t *testing.T) { + router, repo := setupContractRouter(t, "contract-jwt-secret") + + const tunnelFlowGB = int64(500) + const tunnelInFlow = int64(123) + const tunnelOutFlow = int64(456) + const tunnelExpTimeMs = int64(2727251700000) + + now := time.Now().UnixMilli() + res, err := repo.DB().Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + "contract-tunnel", 1.0, 1, "tls", 1, now, now, 1, nil, 0) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, err := res.LastInsertId() + if err != nil { + t.Fatalf("last insert id: %v", err) + } + if _, err := repo.DB().Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, ?, ?, ?, ?, ?)`, + 1, tunnelID, 99999, tunnelFlowGB, tunnelInFlow, tunnelOutFlow, 1, tunnelExpTimeMs, 1); err != nil { + t.Fatalf("insert user_tunnel: %v", err) + } + + t.Run("default user subscription payload", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user", nil) + resp := httptest.NewRecorder() + + router.ServeHTTP(resp, req) + + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatalf("read body: %v", err) + } + + expected := "upload=0; download=0; total=107373108658176; expire=2727251700" + if string(body) != expected { + t.Fatalf("expected body %q, got %q", expected, string(body)) + } + if got := resp.Header().Get("subscription-userinfo"); got != expected { + t.Fatalf("expected subscription-userinfo %q, got %q", expected, got) + } + if !strings.Contains(resp.Header().Get("Content-Type"), "text/plain") { + t.Fatalf("expected text/plain content type, got %q", resp.Header().Get("Content-Type")) + } + }) + + t.Run("tunnel scoped subscription payload", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user&tunnel="+strconv.FormatInt(tunnelID, 10), nil) + resp := httptest.NewRecorder() + + router.ServeHTTP(resp, req) + + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatalf("read body: %v", err) + } + + expected := "upload=123; download=456; total=536870912000; expire=2727251700" + if string(body) != expected { + t.Fatalf("expected body %q, got %q", expected, string(body)) + } + if got := resp.Header().Get("subscription-userinfo"); got != expected { + t.Fatalf("expected subscription-userinfo %q, got %q", expected, got) + } + }) + + t.Run("invalid credentials returns contract error", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=wrong", nil) + resp := httptest.NewRecorder() + + router.ServeHTTP(resp, req) + + assertCodeMsg(t, resp, -1, "鉴权失败") + }) + + t.Run("missing tunnel returns contract error", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user&tunnel=999999", nil) + resp := httptest.NewRecorder() + + router.ServeHTTP(resp, req) + + assertCodeMsg(t, resp, -1, "隧道不存在") + }) +} + +func TestSpeedLimitTunnelsRouteAlias(t *testing.T) { + secret := "contract-jwt-secret" + router, _ := setupContractRouter(t, secret) + + t.Run("missing token blocked", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil) + resp := httptest.NewRecorder() + + router.ServeHTTP(resp, req) + + assertCodeMsg(t, resp, 401, "未登录或token已过期") + }) + + t.Run("admin token receives success envelope", func(t *testing.T) { + token, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil) + req.Header.Set("Authorization", token) + resp := httptest.NewRecorder() + + router.ServeHTTP(resp, req) + + var out response.R + if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code != 0 { + t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg) + } + }) +} + +func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) { + t.Helper() + dbPath := filepath.Join(t.TempDir(), "contract.db") + repo, err := sqlite.Open(dbPath) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + t.Cleanup(func() { + _ = repo.Close() + }) + + h := handler.New(repo, jwtSecret) + return httpserver.NewRouter(h, jwtSecret), repo +}