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"))
+}
+
+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
+}