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