From e2ae241f8c492367ef67a5ec6e1c6bfca9b39dce Mon Sep 17 00:00:00 2001 From: sagit Date: Sat, 7 Feb 2026 13:36:30 +0000 Subject: [PATCH] fix: complete tunnel-create parity with runtime rollback --- .../internal/http/handler/control_plane.go | 11 +- go-backend/internal/http/handler/mutations.go | 614 +++++++++++++++++- .../contract/tunnel_create_contract_test.go | 151 +++++ 3 files changed, 770 insertions(+), 6 deletions(-) create mode 100644 go-backend/tests/contract/tunnel_create_contract_test.go diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 02905d8..345af65 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -44,6 +44,9 @@ type nodeRecord struct { ID int64 Name string ServerIP string + ServerIPv4 string + ServerIPv6 string + Status int PortRange string TCPListenAddr string UDPListenAddr string @@ -192,23 +195,27 @@ func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error) func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) { row := h.repo.DB().QueryRow(` - SELECT id, name, server_ip, port, tcp_listen_addr, udp_listen_addr, interface_name + SELECT id, name, server_ip, server_ip_v4, server_ip_v6, status, port, tcp_listen_addr, udp_listen_addr, interface_name FROM node WHERE id = ? LIMIT 1 `, nodeID) var n nodeRecord + var serverIPv4 sql.NullString + var serverIPv6 sql.NullString var portRange sql.NullString var tcpListen sql.NullString var udpListen sql.NullString var iface sql.NullString - err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &portRange, &tcpListen, &udpListen, &iface) + err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Status, &portRange, &tcpListen, &udpListen, &iface) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, errors.New("节点不存在") } return nil, err } + n.ServerIPv4 = strings.TrimSpace(serverIPv4.String) + n.ServerIPv6 = strings.TrimSpace(serverIPv6.String) n.PortRange = strings.TrimSpace(portRange.String) n.TCPListenAddr = strings.TrimSpace(tcpListen.String) n.UDPListenAddr = strings.TrimSpace(udpListen.String) diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 84f85da..3569756 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -7,6 +7,7 @@ import ( "encoding/json" "errors" "fmt" + "net" "net/http" "sort" "strconv" @@ -525,6 +526,16 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("隧道名称不能为空")) return } + var tunnelNameDup int + if err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE name = ?`, name).Scan(&tunnelNameDup); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if tunnelNameDup > 0 { + response.WriteJSON(w, response.ErrDefault("隧道名称重复")) + return + } + typeVal := asInt(req["type"], 1) flow := asInt64(req["flow"], 1) status := asInt(req["status"], 1) @@ -540,6 +551,15 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { } defer func() { _ = tx.Rollback() }() + runtimeState, err := h.prepareTunnelCreateState(tx, req, typeVal) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + if strings.TrimSpace(inIP) == "" { + inIP = buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes) + } + 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 { @@ -547,6 +567,8 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { return } tunnelID, _ := res.LastInsertId() + runtimeState.TunnelID = tunnelID + applyTunnelPortsToRequest(req, runtimeState) if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -555,6 +577,15 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.Err(-2, err.Error())) return } + if typeVal == 2 { + createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState) + if applyErr != nil { + h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID) + _ = h.deleteTunnelByID(tunnelID) + response.WriteJSON(w, response.ErrDefault(applyErr.Error())) + return + } + } response.WriteJSON(w, response.OKEmpty()) } @@ -1640,7 +1671,566 @@ func queryPairs(db *sql.DB, q string, args ...interface{}) ([][2]int64, error) { return out, rows.Err() } +type tunnelRuntimeNode struct { + NodeID int64 + Protocol string + Strategy string + Inx int + ChainType int + Port int +} + +type tunnelCreateState struct { + TunnelID int64 + Type int + InNodes []tunnelRuntimeNode + ChainHops [][]tunnelRuntimeNode + OutNodes []tunnelRuntimeNode + Nodes map[int64]*nodeRecord + NodeIDList []int64 +} + +func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{}, tunnelType int) (*tunnelCreateState, error) { + state := &tunnelCreateState{ + Type: tunnelType, + InNodes: make([]tunnelRuntimeNode, 0), + ChainHops: make([][]tunnelRuntimeNode, 0), + OutNodes: make([]tunnelRuntimeNode, 0), + Nodes: make(map[int64]*nodeRecord), + } + nodeIDs := make([]int64, 0) + + for _, item := range asMapSlice(req["inNodeId"]) { + nodeID := asInt64(item["nodeId"], 0) + if nodeID <= 0 { + continue + } + nodeIDs = append(nodeIDs, nodeID) + state.InNodes = append(state.InNodes, tunnelRuntimeNode{ + NodeID: nodeID, + Protocol: defaultString(asString(item["protocol"]), "tls"), + Strategy: defaultString(asString(item["strategy"]), "round"), + ChainType: 1, + }) + } + if len(state.InNodes) == 0 { + return nil, errors.New("入口不能为空") + } + + if tunnelType == 2 { + outNodesRaw := asMapSlice(req["outNodeId"]) + if len(outNodesRaw) == 0 { + return nil, errors.New("出口不能为空") + } + + allocated := map[int64]int{} + for _, item := range outNodesRaw { + nodeID := asInt64(item["nodeId"], 0) + if nodeID <= 0 { + continue + } + nodeIDs = append(nodeIDs, nodeID) + port := asInt(item["port"], 0) + if port <= 0 { + var err error + port, err = pickNodePortTx(tx, nodeID, allocated) + if err != nil { + return nil, err + } + } + state.OutNodes = append(state.OutNodes, tunnelRuntimeNode{ + NodeID: nodeID, + Protocol: defaultString(asString(item["protocol"]), "tls"), + Strategy: defaultString(asString(item["strategy"]), "round"), + ChainType: 3, + Port: port, + }) + } + if len(state.OutNodes) == 0 { + return nil, errors.New("出口不能为空") + } + + for hopIdx, hopRaw := range asAnySlice(req["chainNodes"]) { + hop := make([]tunnelRuntimeNode, 0) + for _, item := range asMapSlice(hopRaw) { + nodeID := asInt64(item["nodeId"], 0) + if nodeID <= 0 { + continue + } + nodeIDs = append(nodeIDs, nodeID) + port := asInt(item["port"], 0) + if port <= 0 { + var err error + port, err = pickNodePortTx(tx, nodeID, allocated) + if err != nil { + return nil, err + } + } + hop = append(hop, tunnelRuntimeNode{ + NodeID: nodeID, + Protocol: defaultString(asString(item["protocol"]), "tls"), + Strategy: defaultString(asString(item["strategy"]), "round"), + Inx: hopIdx + 1, + ChainType: 2, + Port: port, + }) + } + if len(hop) > 0 { + state.ChainHops = append(state.ChainHops, hop) + } + } + } + + seen := make(map[int64]struct{}, len(nodeIDs)) + for _, nodeID := range nodeIDs { + if _, ok := seen[nodeID]; ok { + return nil, errors.New("节点重复") + } + seen[nodeID] = struct{}{} + state.NodeIDList = append(state.NodeIDList, nodeID) + node, err := h.getNodeRecord(nodeID) + if err != nil { + if strings.Contains(err.Error(), "不存在") { + return nil, errors.New("节点不存在") + } + return nil, err + } + if node.Status != 1 { + return nil, errors.New("部分节点不在线") + } + state.Nodes[nodeID] = node + } + + return state, nil +} + +func buildTunnelInIP(inNodes []tunnelRuntimeNode, nodes map[int64]*nodeRecord) string { + set := make(map[string]struct{}) + ordered := make([]string, 0) + for _, inNode := range inNodes { + node := nodes[inNode.NodeID] + if node == nil { + continue + } + if v := strings.TrimSpace(node.ServerIPv4); v != "" { + if _, ok := set[v]; !ok { + set[v] = struct{}{} + ordered = append(ordered, v) + } + } + if v := strings.TrimSpace(node.ServerIPv6); v != "" { + if _, ok := set[v]; !ok { + set[v] = struct{}{} + ordered = append(ordered, v) + } + } + if strings.TrimSpace(node.ServerIPv4) == "" && strings.TrimSpace(node.ServerIPv6) == "" { + if v := strings.TrimSpace(node.ServerIP); v != "" { + if _, ok := set[v]; !ok { + set[v] = struct{}{} + ordered = append(ordered, v) + } + } + } + } + return strings.Join(ordered, ",") +} + +func applyTunnelPortsToRequest(req map[string]interface{}, state *tunnelCreateState) { + if req == nil || state == nil { + return + } + outPorts := make(map[int64]int) + for _, n := range state.OutNodes { + outPorts[n.NodeID] = n.Port + } + for _, item := range asMapSlice(req["outNodeId"]) { + nodeID := asInt64(item["nodeId"], 0) + if port, ok := outPorts[nodeID]; ok && port > 0 { + item["port"] = port + } + } + + chainPorts := make(map[int64]int) + for _, hop := range state.ChainHops { + for _, n := range hop { + chainPorts[n.NodeID] = n.Port + } + } + for _, hopRaw := range asAnySlice(req["chainNodes"]) { + for _, item := range asMapSlice(hopRaw) { + nodeID := asInt64(item["nodeId"], 0) + if port, ok := chainPorts[nodeID]; ok && port > 0 { + item["port"] = port + } + } + } +} + +func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64, error) { + if h == nil || state == nil { + return nil, nil, errors.New("invalid tunnel runtime state") + } + createdChains := make([]int64, 0) + createdServices := make([]int64, 0) + if state.Type != 2 { + return createdChains, createdServices, nil + } + + for _, inNode := range state.InNodes { + targets := state.OutNodes + if len(state.ChainHops) > 0 { + targets = state.ChainHops[0] + } + chainData, err := buildTunnelChainConfig(state.TunnelID, inNode.NodeID, targets, state.Nodes) + if err != nil { + return createdChains, createdServices, err + } + if _, err := h.sendNodeCommand(inNode.NodeID, "AddChains", chainData, true, false); err != nil { + return createdChains, createdServices, fmt.Errorf("入口节点 %s 下发转发链失败: %w", nodeDisplayName(state.Nodes[inNode.NodeID]), err) + } + createdChains = append(createdChains, inNode.NodeID) + } + + for i, hop := range state.ChainHops { + nextTargets := state.OutNodes + if i+1 < len(state.ChainHops) { + nextTargets = state.ChainHops[i+1] + } + for _, chainNode := range hop { + chainData, err := buildTunnelChainConfig(state.TunnelID, chainNode.NodeID, nextTargets, state.Nodes) + if err != nil { + return createdChains, createdServices, err + } + if _, err := h.sendNodeCommand(chainNode.NodeID, "AddChains", chainData, true, false); err != nil { + return createdChains, createdServices, fmt.Errorf("转发链节点 %s 下发转发链失败: %w", nodeDisplayName(state.Nodes[chainNode.NodeID]), err) + } + createdChains = append(createdChains, chainNode.NodeID) + + serviceData := buildTunnelChainServiceConfig(state.TunnelID, chainNode, state.Nodes[chainNode.NodeID]) + if _, err := h.sendNodeCommand(chainNode.NodeID, "AddService", serviceData, true, false); err != nil { + return createdChains, createdServices, fmt.Errorf("转发链节点 %s 下发服务失败: %w", nodeDisplayName(state.Nodes[chainNode.NodeID]), err) + } + createdServices = append(createdServices, chainNode.NodeID) + } + } + + for _, outNode := range state.OutNodes { + serviceData := buildTunnelChainServiceConfig(state.TunnelID, outNode, state.Nodes[outNode.NodeID]) + if _, err := h.sendNodeCommand(outNode.NodeID, "AddService", serviceData, true, false); err != nil { + return createdChains, createdServices, fmt.Errorf("出口节点 %s 下发服务失败: %w", nodeDisplayName(state.Nodes[outNode.NodeID]), err) + } + createdServices = append(createdServices, outNode.NodeID) + } + + return createdChains, createdServices, nil +} + +func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tunnelID int64) { + if h == nil || tunnelID <= 0 { + return + } + seenServices := make(map[int64]struct{}) + serviceName := fmt.Sprintf("%d_tls", tunnelID) + for i := len(serviceNodeIDs) - 1; i >= 0; i-- { + nodeID := serviceNodeIDs[i] + if _, ok := seenServices[nodeID]; ok { + continue + } + seenServices[nodeID] = struct{}{} + _, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{serviceName}}, false, true) + } + + seenChains := make(map[int64]struct{}) + chainName := fmt.Sprintf("chains_%d", tunnelID) + for i := len(chainNodeIDs) - 1; i >= 0; i-- { + nodeID := chainNodeIDs[i] + if _, ok := seenChains[nodeID]; ok { + continue + } + seenChains[nodeID] = struct{}{} + _, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true) + } +} + +func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRuntimeNode, nodes map[int64]*nodeRecord) (map[string]interface{}, error) { + fromNode := nodes[fromNodeID] + if fromNode == nil { + return nil, errors.New("节点不存在") + } + if len(targets) == 0 { + return nil, errors.New("转发链目标不能为空") + } + nodeItems := make([]map[string]interface{}, 0, len(targets)) + for idx, target := range targets { + targetNode := nodes[target.NodeID] + if targetNode == nil { + return nil, errors.New("节点不存在") + } + host, err := selectTunnelDialHost(fromNode, targetNode) + if err != nil { + return nil, err + } + port := target.Port + if port <= 0 { + return nil, errors.New("节点端口不能为空") + } + nodeItems = append(nodeItems, map[string]interface{}{ + "name": fmt.Sprintf("node_%d", idx+1), + "addr": processServerAddress(fmt.Sprintf("%s:%d", host, port)), + "connector": map[string]interface{}{ + "type": "relay", + }, + "dialer": map[string]interface{}{ + "type": defaultString(target.Protocol, "tls"), + }, + }) + } + + strategy := defaultString(strings.TrimSpace(targets[0].Strategy), "round") + hop := map[string]interface{}{ + "name": fmt.Sprintf("hop_%d", tunnelID), + "selector": map[string]interface{}{ + "strategy": strategy, + "maxFails": 1, + "failTimeout": int64(600000000000), + }, + "nodes": nodeItems, + } + if strings.TrimSpace(fromNode.InterfaceName) != "" { + hop["interface"] = fromNode.InterfaceName + } + + return map[string]interface{}{ + "name": fmt.Sprintf("chains_%d", tunnelID), + "hops": []map[string]interface{}{hop}, + }, nil +} + +func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode, node *nodeRecord) []map[string]interface{} { + if node == nil { + return nil + } + service := map[string]interface{}{ + "name": fmt.Sprintf("%d_tls", tunnelID), + "addr": fmt.Sprintf("%s:%d", node.TCPListenAddr, chainNode.Port), + "handler": map[string]interface{}{ + "type": "relay", + }, + "listener": map[string]interface{}{ + "type": defaultString(chainNode.Protocol, "tls"), + }, + } + if chainNode.ChainType == 2 { + service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", tunnelID) + } + if chainNode.ChainType == 3 && strings.TrimSpace(node.InterfaceName) != "" { + service["metadata"] = map[string]interface{}{"interface": node.InterfaceName} + } + return []map[string]interface{}{service} +} + +func selectTunnelDialHost(fromNode, toNode *nodeRecord) (string, error) { + if fromNode == nil || toNode == nil { + return "", errors.New("节点不存在") + } + fromV4 := nodeSupportsV4(fromNode) + fromV6 := nodeSupportsV6(fromNode) + toV4 := nodeSupportsV4(toNode) + toV6 := nodeSupportsV6(toNode) + + if fromV4 && toV4 { + host := pickNodeAddressV4(toNode) + if host != "" { + return host, nil + } + } + if fromV6 && toV6 { + host := pickNodeAddressV6(toNode) + if host != "" { + return host, nil + } + } + return "", fmt.Errorf("节点链路不兼容:%s(v4=%t,v6=%t) -> %s(v4=%t,v6=%t)", nodeDisplayName(fromNode), fromV4, fromV6, nodeDisplayName(toNode), toV4, toV6) +} + +func nodeDisplayName(node *nodeRecord) string { + if node == nil { + return "node" + } + if strings.TrimSpace(node.Name) != "" { + return strings.TrimSpace(node.Name) + } + return fmt.Sprintf("node_%d", node.ID) +} + +func nodeSupportsV4(node *nodeRecord) bool { + if node == nil { + return false + } + if strings.TrimSpace(node.ServerIPv4) != "" { + return true + } + if strings.TrimSpace(node.ServerIPv6) != "" { + return false + } + legacy := strings.Trim(strings.TrimSpace(node.ServerIP), "[]") + if legacy == "" { + return false + } + if ip := net.ParseIP(legacy); ip != nil { + return ip.To4() != nil + } + return true +} + +func nodeSupportsV6(node *nodeRecord) bool { + if node == nil { + return false + } + if strings.TrimSpace(node.ServerIPv6) != "" { + return true + } + if strings.TrimSpace(node.ServerIPv4) != "" { + return false + } + legacy := strings.Trim(strings.TrimSpace(node.ServerIP), "[]") + if legacy == "" { + return false + } + if ip := net.ParseIP(legacy); ip != nil { + return ip.To4() == nil + } + return true +} + +func pickNodeAddressV4(node *nodeRecord) string { + if node == nil { + return "" + } + if v := strings.TrimSpace(node.ServerIPv4); v != "" { + return v + } + return strings.TrimSpace(node.ServerIP) +} + +func pickNodeAddressV6(node *nodeRecord) string { + if node == nil { + return "" + } + if v := strings.TrimSpace(node.ServerIPv6); v != "" { + return v + } + return strings.TrimSpace(node.ServerIP) +} + +func pickNodePortTx(tx *sql.Tx, nodeID int64, allocated map[int64]int) (int, error) { + if tx == nil { + return 0, errors.New("database unavailable") + } + if nodeID <= 0 { + return 0, errors.New("节点不存在") + } + if port, ok := allocated[nodeID]; ok && port > 0 { + return port, nil + } + + var portRange string + if err := tx.QueryRow(`SELECT port FROM node WHERE id = ? LIMIT 1`, nodeID).Scan(&portRange); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0, errors.New("节点不存在") + } + return 0, err + } + candidates := parsePortRangeSpec(portRange) + if len(candidates) == 0 { + return 0, errors.New("节点端口已满,无可用端口") + } + + used := map[int]struct{}{} + chainRows, err := tx.Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL`, nodeID) + if err != nil { + return 0, err + } + for chainRows.Next() { + var p sql.NullInt64 + if scanErr := chainRows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 { + used[int(p.Int64)] = struct{}{} + } + } + _ = chainRows.Close() + + forwardRows, err := tx.Query(`SELECT port FROM forward_port WHERE node_id = ?`, nodeID) + if err != nil { + return 0, err + } + for forwardRows.Next() { + var p sql.NullInt64 + if scanErr := forwardRows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 { + used[int(p.Int64)] = struct{}{} + } + } + _ = forwardRows.Close() + + for _, candidate := range candidates { + if candidate <= 0 { + continue + } + if _, ok := used[candidate]; ok { + continue + } + allocated[nodeID] = candidate + return candidate, nil + } + return 0, errors.New("节点端口已满,无可用端口") +} + +func parsePortRangeSpec(input string) []int { + input = strings.TrimSpace(input) + if input == "" { + return nil + } + set := make(map[int]struct{}) + parts := strings.Split(input, ",") + for _, part := range parts { + part = strings.TrimSpace(part) + if part == "" { + continue + } + if strings.Contains(part, "-") { + r := strings.SplitN(part, "-", 2) + if len(r) != 2 { + continue + } + start, err1 := strconv.Atoi(strings.TrimSpace(r[0])) + end, err2 := strconv.Atoi(strings.TrimSpace(r[1])) + if err1 != nil || err2 != nil || start <= 0 || end <= 0 { + continue + } + if end < start { + start, end = end, start + } + for p := start; p <= end; p++ { + set[p] = struct{}{} + } + continue + } + p, err := strconv.Atoi(part) + if err != nil || p <= 0 { + continue + } + set[p] = struct{}{} + } + out := make([]int, 0, len(set)) + for p := range set { + out = append(out, p) + } + sort.Ints(out) + return out +} + func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{}) error { + allocated := map[int64]int{} inNodes := asMapSlice(req["inNodeId"]) for _, n := range inNodes { nodeID := asInt64(n["nodeId"], 0) @@ -1658,8 +2248,16 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{ 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")) + port := asInt(n["port"], 0) + if port <= 0 { + var pickErr error + port, pickErr = pickNodePortTx(tx, nodeID, allocated) + if pickErr != nil { + return pickErr + } + } + _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, ?, NULL, 0, ?)`, + tunnelID, nodeID, port, defaultString(asString(n["protocol"]), "tls")) if err != nil { return err } @@ -1671,8 +2269,16 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{ 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")) + port := asInt(n["port"], 0) + if port <= 0 { + var pickErr error + port, pickErr = pickNodePortTx(tx, nodeID, allocated) + if pickErr != nil { + return pickErr + } + } + _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, ?, ?, ?, ?)`, + tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls")) if err != nil { return err } diff --git a/go-backend/tests/contract/tunnel_create_contract_test.go b/go-backend/tests/contract/tunnel_create_contract_test.go new file mode 100644 index 0000000..407a744 --- /dev/null +++ b/go-backend/tests/contract/tunnel_create_contract_test.go @@ -0,0 +1,151 @@ +package contract_test + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + "time" + + "go-backend/internal/auth" + "go-backend/internal/http/response" +) + +func TestTunnelCreateRuntimeRollbackContract(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + insertNode := func(name, ip, portRange string) int64 { + res, err := repo.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, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) + if err != nil { + t.Fatalf("insert node %s: %v", name, err) + } + id, err := res.LastInsertId() + if err != nil { + t.Fatalf("get node id %s: %v", name, err) + } + return id + } + + entryID := insertNode("create-entry", "10.20.0.1", "30000-30010") + chainID := insertNode("create-chain", "10.20.0.2", "31000-31010") + exitID := insertNode("create-exit", "10.20.0.3", "32000-32010") + + payload := `{"name":"runtime-rollback-tunnel","type":2,"flow":99999,"status":1,"inNodeId":[{"nodeId":` + jsonInt(entryID) + `,"protocol":"tls"}],"chainNodes":[[{"nodeId":` + jsonInt(chainID) + `,"protocol":"tls","strategy":"round"}]],"outNodeId":[{"nodeId":` + jsonInt(exitID) + `,"protocol":"tls"}]}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", bytes.NewBufferString(payload)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + 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 create failure when nodes are offline") + } + if !strings.Contains(out.Msg, "节点") { + t.Fatalf("expected node-related error, got %q", out.Msg) + } + + var tunnelCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE name = ?`, "runtime-rollback-tunnel").Scan(&tunnelCount); err != nil { + t.Fatalf("count tunnel: %v", err) + } + if tunnelCount != 0 { + t.Fatalf("expected tunnel rollback, found %d records", tunnelCount) + } + + var chainCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM chain_tunnel`).Scan(&chainCount); err != nil { + t.Fatalf("count chain_tunnel: %v", err) + } + if chainCount != 0 { + t.Fatalf("expected chain_tunnel rollback, found %d records", chainCount) + } +} + +func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + insertNode := func(name, ip, portRange string) int64 { + res, err := repo.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, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) + if err != nil { + t.Fatalf("insert node %s: %v", name, err) + } + id, err := res.LastInsertId() + if err != nil { + t.Fatalf("get node id %s: %v", name, err) + } + return id + } + + entryID := insertNode("update-entry", "10.30.0.1", "40000-40010") + chainID := insertNode("update-chain", "10.30.0.2", "41000-41010") + exitID := insertNode("update-exit", "10.30.0.3", "42000-42010") + + tunnelRes, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "update-port-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, err := tunnelRes.LastInsertId() + if err != nil { + t.Fatalf("get tunnel id: %v", err) + } + + payload := `{"id":` + jsonInt(tunnelID) + `,"name":"update-port-tunnel","type":2,"flow":99999,"trafficRatio":1.0,"status":1,"inNodeId":[{"nodeId":` + jsonInt(entryID) + `,"protocol":"tls"}],"chainNodes":[[{"nodeId":` + jsonInt(chainID) + `,"protocol":"tls","strategy":"round"}]],"outNodeId":[{"nodeId":` + jsonInt(exitID) + `,"protocol":"tls"}]}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", bytes.NewBufferString(payload)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + assertCode(t, res, 0) + + var chainPort int + if err := repo.DB().QueryRow(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 2 LIMIT 1`, tunnelID).Scan(&chainPort); err != nil { + t.Fatalf("query chain port: %v", err) + } + if chainPort <= 0 { + t.Fatalf("expected chain node port to be assigned, got %d", chainPort) + } + + var outPort int + if err := repo.DB().QueryRow(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 3 LIMIT 1`, tunnelID).Scan(&outPort); err != nil { + t.Fatalf("query out port: %v", err) + } + if outPort <= 0 { + t.Fatalf("expected out node port to be assigned, got %d", outPort) + } +} + +func jsonInt(v int64) string { + return strconv.FormatInt(v, 10) +}