fix: complete tunnel-create parity with runtime rollback

This commit is contained in:
sagit
2026-02-07 13:36:30 +00:00
parent 1a8f424d53
commit e2ae241f8c
3 changed files with 770 additions and 6 deletions
@@ -44,6 +44,9 @@ type nodeRecord struct {
ID int64 ID int64
Name string Name string
ServerIP string ServerIP string
ServerIPv4 string
ServerIPv6 string
Status int
PortRange string PortRange string
TCPListenAddr string TCPListenAddr string
UDPListenAddr string UDPListenAddr string
@@ -192,23 +195,27 @@ func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error)
func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) { func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
row := h.repo.DB().QueryRow(` 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 FROM node
WHERE id = ? WHERE id = ?
LIMIT 1 LIMIT 1
`, nodeID) `, nodeID)
var n nodeRecord var n nodeRecord
var serverIPv4 sql.NullString
var serverIPv6 sql.NullString
var portRange sql.NullString var portRange sql.NullString
var tcpListen sql.NullString var tcpListen sql.NullString
var udpListen sql.NullString var udpListen sql.NullString
var iface 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 err != nil {
if errors.Is(err, sql.ErrNoRows) { if errors.Is(err, sql.ErrNoRows) {
return nil, errors.New("节点不存在") return nil, errors.New("节点不存在")
} }
return nil, err return nil, err
} }
n.ServerIPv4 = strings.TrimSpace(serverIPv4.String)
n.ServerIPv6 = strings.TrimSpace(serverIPv6.String)
n.PortRange = strings.TrimSpace(portRange.String) n.PortRange = strings.TrimSpace(portRange.String)
n.TCPListenAddr = strings.TrimSpace(tcpListen.String) n.TCPListenAddr = strings.TrimSpace(tcpListen.String)
n.UDPListenAddr = strings.TrimSpace(udpListen.String) n.UDPListenAddr = strings.TrimSpace(udpListen.String)
+610 -4
View File
@@ -7,6 +7,7 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"net"
"net/http" "net/http"
"sort" "sort"
"strconv" "strconv"
@@ -525,6 +526,16 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("隧道名称不能为空")) response.WriteJSON(w, response.ErrDefault("隧道名称不能为空"))
return 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) typeVal := asInt(req["type"], 1)
flow := asInt64(req["flow"], 1) flow := asInt64(req["flow"], 1)
status := asInt(req["status"], 1) status := asInt(req["status"], 1)
@@ -540,6 +551,15 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
} }
defer func() { _ = tx.Rollback() }() 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, 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) name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx)
if err != nil { if err != nil {
@@ -547,6 +567,8 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
return return
} }
tunnelID, _ := res.LastInsertId() tunnelID, _ := res.LastInsertId()
runtimeState.TunnelID = tunnelID
applyTunnelPortsToRequest(req, runtimeState)
if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil { if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
@@ -555,6 +577,15 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return 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()) 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() 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 { func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{}) error {
allocated := map[int64]int{}
inNodes := asMapSlice(req["inNodeId"]) inNodes := asMapSlice(req["inNodeId"])
for _, n := range inNodes { for _, n := range inNodes {
nodeID := asInt64(n["nodeId"], 0) nodeID := asInt64(n["nodeId"], 0)
@@ -1658,8 +2248,16 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
if nodeID <= 0 { if nodeID <= 0 {
continue continue
} }
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, NULL, NULL, 0, ?)`, port := asInt(n["port"], 0)
tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls")) 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 { if err != nil {
return err return err
} }
@@ -1671,8 +2269,16 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
if nodeID <= 0 { if nodeID <= 0 {
continue continue
} }
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, NULL, ?, ?, ?)`, port := asInt(n["port"], 0)
tunnelID, nodeID, defaultString(asString(n["strategy"]), "round"), i, defaultString(asString(n["protocol"]), "tls")) 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 { if err != nil {
return err return err
} }
@@ -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)
}