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
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)
+610 -4
View File
@@ -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
}
@@ -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)
}