mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-29 07:56:37 +08:00
fix: complete tunnel-create parity with runtime rollback
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user