mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-08 10:46:37 +08:00
fix: complete diagnosis parity and tunnel visibility on Go backend
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"sort"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -43,6 +44,7 @@ type nodeRecord struct {
|
|||||||
ID int64
|
ID int64
|
||||||
Name string
|
Name string
|
||||||
ServerIP string
|
ServerIP string
|
||||||
|
PortRange string
|
||||||
TCPListenAddr string
|
TCPListenAddr string
|
||||||
UDPListenAddr string
|
UDPListenAddr string
|
||||||
InterfaceName string
|
InterfaceName string
|
||||||
@@ -52,9 +54,16 @@ type chainNodeRecord struct {
|
|||||||
ChainType int
|
ChainType int
|
||||||
Inx int64
|
Inx int64
|
||||||
NodeID int64
|
NodeID int64
|
||||||
|
Port int
|
||||||
NodeName string
|
NodeName string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type diagnosisTarget struct {
|
||||||
|
Address string
|
||||||
|
IP string
|
||||||
|
Port int
|
||||||
|
}
|
||||||
|
|
||||||
func (h *Handler) resolveForwardAccess(r *http.Request, forwardID int64) (*forwardRecord, int64, int, error) {
|
func (h *Handler) resolveForwardAccess(r *http.Request, forwardID int64) (*forwardRecord, int64, int, error) {
|
||||||
userID, roleID, err := userRoleFromRequest(r)
|
userID, roleID, err := userRoleFromRequest(r)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -183,22 +192,24 @@ 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, tcp_listen_addr, udp_listen_addr, interface_name
|
SELECT id, name, server_ip, 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 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, &tcpListen, &udpListen, &iface)
|
err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &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.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)
|
||||||
n.InterfaceName = strings.TrimSpace(iface.String)
|
n.InterfaceName = strings.TrimSpace(iface.String)
|
||||||
@@ -343,68 +354,102 @@ func (h *Handler) diagnoseForwardRuntime(forward *forwardRecord) (map[string]int
|
|||||||
if forward == nil {
|
if forward == nil {
|
||||||
return nil, errForwardNotFound
|
return nil, errForwardNotFound
|
||||||
}
|
}
|
||||||
targets := splitRemoteTargets(forward.RemoteAddr)
|
targets, err := resolveDiagnosisTargets(forward.RemoteAddr)
|
||||||
if len(targets) == 0 {
|
|
||||||
return nil, errors.New("目标地址不能为空")
|
|
||||||
}
|
|
||||||
ip, port, err := parseTargetAddress(targets[0])
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("目标地址格式错误")
|
|
||||||
}
|
|
||||||
|
|
||||||
ports, err := h.listForwardPorts(forward.ID)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if len(ports) == 0 {
|
|
||||||
return nil, errors.New("转发入口端口不存在")
|
tunnel, err := h.getTunnelRecord(forward.TunnelID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
results := make([]map[string]interface{}, 0, len(ports))
|
chainRows, err := h.listChainNodesForTunnel(forward.TunnelID)
|
||||||
seen := map[int64]struct{}{}
|
if err != nil {
|
||||||
for _, fp := range ports {
|
return nil, err
|
||||||
if _, ok := seen[fp.NodeID]; ok {
|
}
|
||||||
continue
|
if len(chainRows) == 0 {
|
||||||
}
|
return nil, errors.New("隧道配置不完整")
|
||||||
seen[fp.NodeID] = struct{}{}
|
}
|
||||||
node, nodeErr := h.getNodeRecord(fp.NodeID)
|
|
||||||
nodeName := fmt.Sprintf("node_%d", fp.NodeID)
|
|
||||||
if nodeErr == nil {
|
|
||||||
nodeName = node.Name
|
|
||||||
}
|
|
||||||
|
|
||||||
resultItem := map[string]interface{}{
|
inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
|
||||||
"nodeName": nodeName,
|
results := make([]map[string]interface{}, 0, len(chainRows)*2+len(targets))
|
||||||
"nodeId": strconv.FormatInt(fp.NodeID, 10),
|
nodeCache := map[int64]*nodeRecord{}
|
||||||
"targetIp": ip,
|
|
||||||
"targetPort": port,
|
|
||||||
"averageTime": 0,
|
|
||||||
"packetLoss": 100,
|
|
||||||
}
|
|
||||||
|
|
||||||
if nodeErr != nil {
|
switch tunnel.Type {
|
||||||
resultItem["success"] = false
|
case 1:
|
||||||
resultItem["description"] = "节点信息读取失败"
|
for _, inNode := range inNodes {
|
||||||
resultItem["errorMessage"] = nodeErr.Error()
|
for _, target := range targets {
|
||||||
results = append(results, resultItem)
|
description := fmt.Sprintf("入口(%s)->目标(%s)", inNode.NodeName, target.Address)
|
||||||
continue
|
h.appendPathDiagnosis(&results, nodeCache, inNode.NodeID, target.IP, target.Port, description, map[string]interface{}{
|
||||||
}
|
"fromChainType": 1,
|
||||||
|
})
|
||||||
pingData, pingErr := h.tcpPingViaNode(fp.NodeID, ip, port)
|
}
|
||||||
if pingErr != nil {
|
}
|
||||||
resultItem["success"] = false
|
case 2:
|
||||||
resultItem["description"] = "诊断失败"
|
for _, inNode := range inNodes {
|
||||||
resultItem["errorMessage"] = pingErr.Error()
|
if len(chainHops) > 0 {
|
||||||
} else {
|
for _, firstNode := range chainHops[0] {
|
||||||
resultItem["success"] = asBool(pingData["success"], false)
|
description := fmt.Sprintf("入口(%s)->第1跳(%s)", inNode.NodeName, firstNode.NodeName)
|
||||||
resultItem["description"] = "诊断完成"
|
h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, firstNode, description, map[string]interface{}{
|
||||||
resultItem["averageTime"] = asFloat(pingData["averageTime"], 0)
|
"fromChainType": 1,
|
||||||
resultItem["packetLoss"] = asFloat(pingData["packetLoss"], 100)
|
"toChainType": 2,
|
||||||
if msg := asString(pingData["errorMessage"]); msg != "" {
|
"toInx": firstNode.Inx,
|
||||||
resultItem["errorMessage"] = msg
|
})
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
for _, outNode := range outNodes {
|
||||||
|
description := fmt.Sprintf("入口(%s)->出口(%s)", inNode.NodeName, outNode.NodeName)
|
||||||
|
h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, outNode, description, map[string]interface{}{
|
||||||
|
"fromChainType": 1,
|
||||||
|
"toChainType": 3,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, hop := range chainHops {
|
||||||
|
for _, currentNode := range hop {
|
||||||
|
if i+1 < len(chainHops) {
|
||||||
|
for _, nextNode := range chainHops[i+1] {
|
||||||
|
description := fmt.Sprintf("第%d跳(%s)->第%d跳(%s)", i+1, currentNode.NodeName, i+2, nextNode.NodeName)
|
||||||
|
h.appendChainHopDiagnosis(&results, nodeCache, currentNode.NodeID, nextNode, description, map[string]interface{}{
|
||||||
|
"fromChainType": 2,
|
||||||
|
"fromInx": currentNode.Inx,
|
||||||
|
"toChainType": 2,
|
||||||
|
"toInx": nextNode.Inx,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
for _, outNode := range outNodes {
|
||||||
|
description := fmt.Sprintf("第%d跳(%s)->出口(%s)", i+1, currentNode.NodeName, outNode.NodeName)
|
||||||
|
h.appendChainHopDiagnosis(&results, nodeCache, currentNode.NodeID, outNode, description, map[string]interface{}{
|
||||||
|
"fromChainType": 2,
|
||||||
|
"fromInx": currentNode.Inx,
|
||||||
|
"toChainType": 3,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, outNode := range outNodes {
|
||||||
|
for _, target := range targets {
|
||||||
|
description := fmt.Sprintf("出口(%s)->目标(%s)", outNode.NodeName, target.Address)
|
||||||
|
h.appendPathDiagnosis(&results, nodeCache, outNode.NodeID, target.IP, target.Port, description, map[string]interface{}{
|
||||||
|
"fromChainType": 3,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
for _, inNode := range inNodes {
|
||||||
|
for _, target := range targets {
|
||||||
|
description := fmt.Sprintf("入口(%s)->目标(%s)", inNode.NodeName, target.Address)
|
||||||
|
h.appendPathDiagnosis(&results, nodeCache, inNode.NodeID, target.IP, target.Port, description, map[string]interface{}{
|
||||||
|
"fromChainType": 1,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
results = append(results, resultItem)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
payload := map[string]interface{}{
|
payload := map[string]interface{}{
|
||||||
@@ -437,53 +482,78 @@ func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{},
|
|||||||
return nil, errors.New("隧道配置不完整")
|
return nil, errors.New("隧道配置不完整")
|
||||||
}
|
}
|
||||||
|
|
||||||
results := make([]map[string]interface{}, 0, len(chainRows))
|
inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
|
||||||
seen := map[int64]struct{}{}
|
results := make([]map[string]interface{}, 0, len(chainRows)*2)
|
||||||
for _, row := range chainRows {
|
nodeCache := map[int64]*nodeRecord{}
|
||||||
if _, ok := seen[row.NodeID]; ok {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
seen[row.NodeID] = struct{}{}
|
|
||||||
|
|
||||||
desc := fmt.Sprintf("节点(%s)->外网", row.NodeName)
|
switch tunnel.Type {
|
||||||
switch row.ChainType {
|
case 1:
|
||||||
case 1:
|
for _, inNode := range inNodes {
|
||||||
desc = fmt.Sprintf("入口(%s)->外网", row.NodeName)
|
description := fmt.Sprintf("入口(%s)->外网", inNode.NodeName)
|
||||||
case 2:
|
h.appendPathDiagnosis(&results, nodeCache, inNode.NodeID, "www.google.com", 443, description, map[string]interface{}{
|
||||||
desc = fmt.Sprintf("中继(%s)->外网", row.NodeName)
|
"fromChainType": 1,
|
||||||
case 3:
|
})
|
||||||
desc = fmt.Sprintf("出口(%s)->外网", row.NodeName)
|
|
||||||
}
|
}
|
||||||
|
case 2:
|
||||||
resultItem := map[string]interface{}{
|
for _, inNode := range inNodes {
|
||||||
"nodeName": row.NodeName,
|
if len(chainHops) > 0 {
|
||||||
"nodeId": strconv.FormatInt(row.NodeID, 10),
|
for _, firstNode := range chainHops[0] {
|
||||||
"targetIp": "www.google.com",
|
description := fmt.Sprintf("入口(%s)->第1跳(%s)", inNode.NodeName, firstNode.NodeName)
|
||||||
"targetPort": 443,
|
h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, firstNode, description, map[string]interface{}{
|
||||||
"averageTime": 0,
|
"fromChainType": 1,
|
||||||
"packetLoss": 100,
|
"toChainType": 2,
|
||||||
"fromChainType": row.ChainType,
|
"toInx": firstNode.Inx,
|
||||||
"description": desc,
|
})
|
||||||
}
|
}
|
||||||
if row.ChainType == 2 {
|
} else {
|
||||||
resultItem["fromInx"] = row.Inx
|
for _, outNode := range outNodes {
|
||||||
}
|
description := fmt.Sprintf("入口(%s)->出口(%s)", inNode.NodeName, outNode.NodeName)
|
||||||
|
h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, outNode, description, map[string]interface{}{
|
||||||
pingData, pingErr := h.tcpPingViaNode(row.NodeID, "www.google.com", 443)
|
"fromChainType": 1,
|
||||||
if pingErr != nil {
|
"toChainType": 3,
|
||||||
resultItem["success"] = false
|
})
|
||||||
resultItem["description"] = desc + " 连通性检查失败"
|
}
|
||||||
resultItem["errorMessage"] = pingErr.Error()
|
|
||||||
} else {
|
|
||||||
resultItem["success"] = asBool(pingData["success"], false)
|
|
||||||
resultItem["description"] = desc + " 连通性检查完成"
|
|
||||||
resultItem["averageTime"] = asFloat(pingData["averageTime"], 0)
|
|
||||||
resultItem["packetLoss"] = asFloat(pingData["packetLoss"], 100)
|
|
||||||
if msg := asString(pingData["errorMessage"]); msg != "" {
|
|
||||||
resultItem["errorMessage"] = msg
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
results = append(results, resultItem)
|
|
||||||
|
for i, hop := range chainHops {
|
||||||
|
for _, currentNode := range hop {
|
||||||
|
if i+1 < len(chainHops) {
|
||||||
|
for _, nextNode := range chainHops[i+1] {
|
||||||
|
description := fmt.Sprintf("第%d跳(%s)->第%d跳(%s)", i+1, currentNode.NodeName, i+2, nextNode.NodeName)
|
||||||
|
h.appendChainHopDiagnosis(&results, nodeCache, currentNode.NodeID, nextNode, description, map[string]interface{}{
|
||||||
|
"fromChainType": 2,
|
||||||
|
"fromInx": currentNode.Inx,
|
||||||
|
"toChainType": 2,
|
||||||
|
"toInx": nextNode.Inx,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
for _, outNode := range outNodes {
|
||||||
|
description := fmt.Sprintf("第%d跳(%s)->出口(%s)", i+1, currentNode.NodeName, outNode.NodeName)
|
||||||
|
h.appendChainHopDiagnosis(&results, nodeCache, currentNode.NodeID, outNode, description, map[string]interface{}{
|
||||||
|
"fromChainType": 2,
|
||||||
|
"fromInx": currentNode.Inx,
|
||||||
|
"toChainType": 3,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, outNode := range outNodes {
|
||||||
|
description := fmt.Sprintf("出口(%s)->外网", outNode.NodeName)
|
||||||
|
h.appendPathDiagnosis(&results, nodeCache, outNode.NodeID, "www.google.com", 443, description, map[string]interface{}{
|
||||||
|
"fromChainType": 3,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
for _, inNode := range inNodes {
|
||||||
|
description := fmt.Sprintf("入口(%s)->外网", inNode.NodeName)
|
||||||
|
h.appendPathDiagnosis(&results, nodeCache, inNode.NodeID, "www.google.com", 443, description, map[string]interface{}{
|
||||||
|
"fromChainType": 1,
|
||||||
|
})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
payload := map[string]interface{}{
|
payload := map[string]interface{}{
|
||||||
@@ -495,9 +565,198 @@ func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{},
|
|||||||
return payload, nil
|
return payload, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func splitChainNodeGroups(rows []chainNodeRecord) ([]chainNodeRecord, [][]chainNodeRecord, []chainNodeRecord) {
|
||||||
|
inNodes := make([]chainNodeRecord, 0)
|
||||||
|
outNodes := make([]chainNodeRecord, 0)
|
||||||
|
chainByInx := map[int64][]chainNodeRecord{}
|
||||||
|
hopOrder := make([]int64, 0)
|
||||||
|
|
||||||
|
for _, row := range rows {
|
||||||
|
switch row.ChainType {
|
||||||
|
case 1:
|
||||||
|
inNodes = append(inNodes, row)
|
||||||
|
case 2:
|
||||||
|
if _, ok := chainByInx[row.Inx]; !ok {
|
||||||
|
hopOrder = append(hopOrder, row.Inx)
|
||||||
|
}
|
||||||
|
chainByInx[row.Inx] = append(chainByInx[row.Inx], row)
|
||||||
|
case 3:
|
||||||
|
outNodes = append(outNodes, row)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sort.Slice(hopOrder, func(i, j int) bool { return hopOrder[i] < hopOrder[j] })
|
||||||
|
chainHops := make([][]chainNodeRecord, 0, len(hopOrder))
|
||||||
|
for _, inx := range hopOrder {
|
||||||
|
chainHops = append(chainHops, chainByInx[inx])
|
||||||
|
}
|
||||||
|
|
||||||
|
return inNodes, chainHops, outNodes
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveDiagnosisTargets(remoteAddr string) ([]diagnosisTarget, error) {
|
||||||
|
rawTargets := splitRemoteTargets(remoteAddr)
|
||||||
|
if len(rawTargets) == 0 {
|
||||||
|
return nil, errors.New("目标地址不能为空")
|
||||||
|
}
|
||||||
|
|
||||||
|
targets := make([]diagnosisTarget, 0, len(rawTargets))
|
||||||
|
for _, raw := range rawTargets {
|
||||||
|
ip, port, err := parseTargetAddress(raw)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
targets = append(targets, diagnosisTarget{Address: raw, IP: ip, Port: port})
|
||||||
|
}
|
||||||
|
if len(targets) == 0 {
|
||||||
|
return nil, errors.New("目标地址格式错误")
|
||||||
|
}
|
||||||
|
return targets, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) cachedNode(nodeCache map[int64]*nodeRecord, nodeID int64) (*nodeRecord, error) {
|
||||||
|
if node, ok := nodeCache[nodeID]; ok {
|
||||||
|
return node, nil
|
||||||
|
}
|
||||||
|
node, err := h.getNodeRecord(nodeID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
nodeCache[nodeID] = node
|
||||||
|
return node, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func newDiagnosisResultItem(fromNodeID int64, targetIP string, targetPort int, description string, metadata map[string]interface{}) map[string]interface{} {
|
||||||
|
item := map[string]interface{}{
|
||||||
|
"nodeName": fmt.Sprintf("node_%d", fromNodeID),
|
||||||
|
"nodeId": strconv.FormatInt(fromNodeID, 10),
|
||||||
|
"targetIp": targetIP,
|
||||||
|
"targetPort": targetPort,
|
||||||
|
"description": description,
|
||||||
|
"averageTime": 0,
|
||||||
|
"packetLoss": 100,
|
||||||
|
}
|
||||||
|
for k, v := range metadata {
|
||||||
|
item[k] = v
|
||||||
|
}
|
||||||
|
return item
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) appendFailedDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, targetIP string, targetPort int, description string, metadata map[string]interface{}, message string) {
|
||||||
|
item := newDiagnosisResultItem(fromNodeID, targetIP, targetPort, description, metadata)
|
||||||
|
if node, err := h.cachedNode(nodeCache, fromNodeID); err == nil {
|
||||||
|
item["nodeName"] = node.Name
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(message) == "" {
|
||||||
|
message = "TCP连接失败"
|
||||||
|
}
|
||||||
|
item["success"] = false
|
||||||
|
item["message"] = message
|
||||||
|
*results = append(*results, item)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) appendPathDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, targetIP string, targetPort int, description string, metadata map[string]interface{}) {
|
||||||
|
item := newDiagnosisResultItem(fromNodeID, targetIP, targetPort, description, metadata)
|
||||||
|
|
||||||
|
fromNode, err := h.cachedNode(nodeCache, fromNodeID)
|
||||||
|
if err != nil {
|
||||||
|
item["success"] = false
|
||||||
|
item["message"] = err.Error()
|
||||||
|
*results = append(*results, item)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
item["nodeName"] = fromNode.Name
|
||||||
|
|
||||||
|
pingData, pingErr := h.tcpPingViaNode(fromNodeID, targetIP, targetPort)
|
||||||
|
if pingErr != nil {
|
||||||
|
item["success"] = false
|
||||||
|
item["message"] = pingErr.Error()
|
||||||
|
*results = append(*results, item)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
success := asBool(pingData["success"], false)
|
||||||
|
item["success"] = success
|
||||||
|
item["averageTime"] = asFloat(pingData["averageTime"], 0)
|
||||||
|
item["packetLoss"] = asFloat(pingData["packetLoss"], 100)
|
||||||
|
|
||||||
|
message := strings.TrimSpace(asString(pingData["message"]))
|
||||||
|
if success {
|
||||||
|
if message == "" {
|
||||||
|
message = "TCP连接成功"
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if message == "" {
|
||||||
|
message = strings.TrimSpace(asString(pingData["errorMessage"]))
|
||||||
|
}
|
||||||
|
if message == "" {
|
||||||
|
message = "TCP连接失败"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
item["message"] = message
|
||||||
|
*results = append(*results, item)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) appendChainHopDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, toNode chainNodeRecord, description string, metadata map[string]interface{}) {
|
||||||
|
targetNode, err := h.cachedNode(nodeCache, toNode.NodeID)
|
||||||
|
if err != nil {
|
||||||
|
h.appendFailedDiagnosis(results, nodeCache, fromNodeID, "", 0, description, metadata, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
targetIP, targetPort, err := resolveChainProbeTarget(targetNode, toNode.Port)
|
||||||
|
if err != nil {
|
||||||
|
h.appendFailedDiagnosis(results, nodeCache, fromNodeID, strings.Trim(strings.TrimSpace(targetNode.ServerIP), "[]"), toNode.Port, description, metadata, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
h.appendPathDiagnosis(results, nodeCache, fromNodeID, targetIP, targetPort, description, metadata)
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveChainProbeTarget(targetNode *nodeRecord, preferredPort int) (string, int, error) {
|
||||||
|
if targetNode == nil {
|
||||||
|
return "", 0, errors.New("目标节点不存在")
|
||||||
|
}
|
||||||
|
host := strings.Trim(strings.TrimSpace(targetNode.ServerIP), "[]")
|
||||||
|
if host == "" {
|
||||||
|
return "", 0, errors.New("目标节点地址为空")
|
||||||
|
}
|
||||||
|
port := preferredPort
|
||||||
|
if port <= 0 {
|
||||||
|
port = firstPortFromRange(targetNode.PortRange)
|
||||||
|
}
|
||||||
|
if port <= 0 {
|
||||||
|
port = 443
|
||||||
|
}
|
||||||
|
return host, port, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func firstPortFromRange(portRange string) int {
|
||||||
|
portRange = strings.TrimSpace(portRange)
|
||||||
|
if portRange == "" {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
first := strings.Split(portRange, ",")[0]
|
||||||
|
first = strings.TrimSpace(first)
|
||||||
|
if strings.Contains(first, "-") {
|
||||||
|
parts := strings.SplitN(first, "-", 2)
|
||||||
|
if len(parts) != 2 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
p, err := strconv.Atoi(strings.TrimSpace(parts[0]))
|
||||||
|
if err != nil || p <= 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
p, err := strconv.Atoi(first)
|
||||||
|
if err != nil || p <= 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) {
|
func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) {
|
||||||
rows, err := h.repo.DB().Query(`
|
rows, err := h.repo.DB().Query(`
|
||||||
SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, n.name
|
SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name
|
||||||
FROM chain_tunnel ct
|
FROM chain_tunnel ct
|
||||||
LEFT JOIN node n ON n.id = ct.node_id
|
LEFT JOIN node n ON n.id = ct.node_id
|
||||||
WHERE ct.tunnel_id = ?
|
WHERE ct.tunnel_id = ?
|
||||||
@@ -512,7 +771,7 @@ func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, er
|
|||||||
for rows.Next() {
|
for rows.Next() {
|
||||||
var item chainNodeRecord
|
var item chainNodeRecord
|
||||||
var name sql.NullString
|
var name sql.NullString
|
||||||
if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &name); err != nil {
|
if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &item.Port, &name); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(name.String) == "" {
|
if strings.TrimSpace(name.String) == "" {
|
||||||
|
|||||||
@@ -443,13 +443,18 @@ func (h *Handler) userTunnelVisibleList(w http.ResponseWriter, r *http.Request)
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userID, err := userIDFromRequest(r)
|
userID, roleID, err := userRoleFromRequest(r)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
items, err := h.repo.ListUserAccessibleTunnels(userID)
|
items := make([]map[string]interface{}, 0)
|
||||||
|
if roleID == 0 {
|
||||||
|
items, err = h.repo.ListEnabledTunnelSummaries()
|
||||||
|
} else {
|
||||||
|
items, err = h.repo.ListUserAccessibleTunnels(userID)
|
||||||
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -656,10 +656,10 @@ func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]inter
|
|||||||
}
|
}
|
||||||
|
|
||||||
rows, err := r.db.Query(`
|
rows, err := r.db.Query(`
|
||||||
SELECT t.id, t.name
|
SELECT DISTINCT t.id, t.name
|
||||||
FROM user_tunnel ut
|
FROM user_tunnel ut
|
||||||
JOIN tunnel t ON t.id = ut.tunnel_id
|
JOIN tunnel t ON t.id = ut.tunnel_id
|
||||||
WHERE ut.user_id = ? AND ut.status = 1
|
WHERE ut.user_id = ? AND t.status = 1
|
||||||
ORDER BY t.inx ASC, t.id ASC
|
ORDER BY t.inx ASC, t.id ASC
|
||||||
`, userID)
|
`, userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -683,6 +683,38 @@ func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]inter
|
|||||||
return items, nil
|
return items, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *Repository) ListEnabledTunnelSummaries() ([]map[string]interface{}, error) {
|
||||||
|
if r == nil || r.db == nil {
|
||||||
|
return nil, errors.New("repository not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := r.db.Query(`
|
||||||
|
SELECT id, name
|
||||||
|
FROM tunnel
|
||||||
|
WHERE status = 1
|
||||||
|
ORDER BY inx ASC, id ASC
|
||||||
|
`)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
items := make([]map[string]interface{}, 0)
|
||||||
|
for rows.Next() {
|
||||||
|
var id int64
|
||||||
|
var name string
|
||||||
|
if err := rows.Scan(&id, &name); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
items = append(items, map[string]interface{}{"id": id, "name": name})
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return items, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||||
if r == nil || r.db == nil {
|
if r == nil || r.db == nil {
|
||||||
return nil, errors.New("repository not initialized")
|
return nil, errors.New("repository not initialized")
|
||||||
|
|||||||
@@ -0,0 +1,239 @@
|
|||||||
|
package contract
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"go-backend/internal/auth"
|
||||||
|
httpserver "go-backend/internal/http"
|
||||||
|
"go-backend/internal/http/handler"
|
||||||
|
"go-backend/internal/http/response"
|
||||||
|
"go-backend/internal/store/sqlite"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDiagnosisChainCoverageContracts(t *testing.T) {
|
||||||
|
secret := "contract-jwt-secret"
|
||||||
|
router, repo := setupDiagnosisContractRouter(t, secret)
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||||
|
VALUES(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||||
|
`, now, now); err != nil {
|
||||||
|
t.Fatalf("insert user: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tunnelRes, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
`, "diagnose-chain-tunnel", 1.0, 2, "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)
|
||||||
|
}
|
||||||
|
|
||||||
|
insertNode := func(name, ip 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, "", "30000-30010", "", "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
|
||||||
|
}
|
||||||
|
|
||||||
|
entryNodeID := insertNode("entry-node", "10.0.1.10")
|
||||||
|
chainNodeID := insertNode("chain-node", "10.0.1.20")
|
||||||
|
exitNodeID := insertNode("exit-node", "10.0.1.30")
|
||||||
|
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||||
|
VALUES(?, 1, ?, 30001, 'round', 1, 'tls')
|
||||||
|
`, tunnelID, entryNodeID); err != nil {
|
||||||
|
t.Fatalf("insert entry chain: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||||
|
VALUES(?, 2, ?, 30002, 'round', 1, 'tls')
|
||||||
|
`, tunnelID, chainNodeID); err != nil {
|
||||||
|
t.Fatalf("insert middle chain: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||||
|
VALUES(?, 3, ?, 30003, 'round', 1, 'tls')
|
||||||
|
`, tunnelID, exitNodeID); err != nil {
|
||||||
|
t.Fatalf("insert exit chain: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
forwardRes, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
||||||
|
`, 2, "normal_user", "chain-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("insert forward: %v", err)
|
||||||
|
}
|
||||||
|
forwardID, err := forwardRes.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get forward id: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate user token: %v", err)
|
||||||
|
}
|
||||||
|
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate admin token: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("forward diagnose includes entry chain exit paths", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/diagnose", bytes.NewBufferString(`{"forwardId":`+strconv.FormatInt(forwardID, 10)+`}`))
|
||||||
|
req.Header.Set("Authorization", userToken)
|
||||||
|
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 code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
payload, ok := out.Data.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected object payload, got %T", out.Data)
|
||||||
|
}
|
||||||
|
results, ok := payload["results"].([]interface{})
|
||||||
|
if !ok || len(results) == 0 {
|
||||||
|
t.Fatalf("expected non-empty results, got %v", payload["results"])
|
||||||
|
}
|
||||||
|
|
||||||
|
hasEntryToChain := false
|
||||||
|
hasChainToExit := false
|
||||||
|
hasExitToTarget := false
|
||||||
|
for _, raw := range results {
|
||||||
|
item, ok := raw.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected result object, got %T", raw)
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(valueAsString(item["message"])) == "" {
|
||||||
|
t.Fatalf("expected non-empty message field")
|
||||||
|
}
|
||||||
|
from := valueAsInt(item["fromChainType"])
|
||||||
|
to := valueAsInt(item["toChainType"])
|
||||||
|
if from == 1 && to == 2 {
|
||||||
|
hasEntryToChain = true
|
||||||
|
}
|
||||||
|
if from == 2 && to == 3 {
|
||||||
|
hasChainToExit = true
|
||||||
|
}
|
||||||
|
if from == 3 {
|
||||||
|
hasExitToTarget = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hasEntryToChain || !hasChainToExit || !hasExitToTarget {
|
||||||
|
t.Fatalf("expected entry->chain, chain->exit, exit->target coverage; got entry=%v chain=%v exit=%v", hasEntryToChain, hasChainToExit, hasExitToTarget)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("tunnel diagnose includes entry chain exit groups", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/diagnose", bytes.NewBufferString(`{"tunnelId":`+strconv.FormatInt(tunnelID, 10)+`}`))
|
||||||
|
req.Header.Set("Authorization", adminToken)
|
||||||
|
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 code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
payload, ok := out.Data.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected object payload, got %T", out.Data)
|
||||||
|
}
|
||||||
|
results, ok := payload["results"].([]interface{})
|
||||||
|
if !ok || len(results) == 0 {
|
||||||
|
t.Fatalf("expected non-empty results, got %v", payload["results"])
|
||||||
|
}
|
||||||
|
|
||||||
|
hasEntry := false
|
||||||
|
hasChain := false
|
||||||
|
hasExit := false
|
||||||
|
for _, raw := range results {
|
||||||
|
item, ok := raw.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected result object, got %T", raw)
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(valueAsString(item["message"])) == "" {
|
||||||
|
t.Fatalf("expected non-empty message field")
|
||||||
|
}
|
||||||
|
switch valueAsInt(item["fromChainType"]) {
|
||||||
|
case 1:
|
||||||
|
hasEntry = true
|
||||||
|
case 2:
|
||||||
|
hasChain = true
|
||||||
|
case 3:
|
||||||
|
hasExit = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hasEntry || !hasChain || !hasExit {
|
||||||
|
t.Fatalf("expected entry/chain/exit groups, got entry=%v chain=%v exit=%v", hasEntry, hasChain, hasExit)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func valueAsInt(v interface{}) int {
|
||||||
|
switch n := v.(type) {
|
||||||
|
case float64:
|
||||||
|
return int(n)
|
||||||
|
case int:
|
||||||
|
return n
|
||||||
|
case int64:
|
||||||
|
return int(n)
|
||||||
|
default:
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func valueAsString(v interface{}) string {
|
||||||
|
s, _ := v.(string)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupDiagnosisContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) {
|
||||||
|
t.Helper()
|
||||||
|
dbPath := filepath.Join(t.TempDir(), "diagnosis-contract.db")
|
||||||
|
repo, err := sqlite.Open(dbPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open sqlite: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_ = repo.Close()
|
||||||
|
})
|
||||||
|
|
||||||
|
h := handler.New(repo, jwtSecret)
|
||||||
|
return httpserver.NewRouter(h, jwtSecret), repo
|
||||||
|
}
|
||||||
@@ -37,6 +37,25 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
|
|||||||
t.Fatalf("get tunnel id: %v", err)
|
t.Fatalf("get tunnel id: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
nodeRes, 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
`, "entry-node", "entry-secret", "10.0.0.10", "10.0.0.10", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("insert node: %v", err)
|
||||||
|
}
|
||||||
|
entryNodeID, err := nodeRes.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get node id: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||||
|
VALUES(?, 1, ?, 20001, 'round', 1, 'tls')
|
||||||
|
`, tunnelID, entryNodeID); err != nil {
|
||||||
|
t.Fatalf("insert chain_tunnel: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
resAdmin, err := repo.DB().Exec(`
|
resAdmin, err := repo.DB().Exec(`
|
||||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
||||||
@@ -65,6 +84,10 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("generate user token: %v", err)
|
t.Fatalf("generate user token: %v", err)
|
||||||
}
|
}
|
||||||
|
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate admin token: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
t.Run("non-owner cannot delete another user's forward", func(t *testing.T) {
|
t.Run("non-owner cannot delete another user's forward", func(t *testing.T) {
|
||||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/delete", bytes.NewBufferString(`{"id":`+jsonNumber(adminForwardID)+`}`))
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/delete", bytes.NewBufferString(`{"id":`+jsonNumber(adminForwardID)+`}`))
|
||||||
@@ -106,7 +129,7 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("diagnose no longer returns hardcoded success", func(t *testing.T) {
|
t.Run("forward diagnose returns structured payload", func(t *testing.T) {
|
||||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/diagnose", bytes.NewBufferString(`{"forwardId":`+jsonNumber(userForwardID)+`}`))
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/diagnose", bytes.NewBufferString(`{"forwardId":`+jsonNumber(userForwardID)+`}`))
|
||||||
req.Header.Set("Authorization", userToken)
|
req.Header.Set("Authorization", userToken)
|
||||||
res := httptest.NewRecorder()
|
res := httptest.NewRecorder()
|
||||||
@@ -117,8 +140,59 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
|
|||||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||||
t.Fatalf("decode response: %v", err)
|
t.Fatalf("decode response: %v", err)
|
||||||
}
|
}
|
||||||
if out.Code == 0 {
|
if out.Code != 0 {
|
||||||
t.Fatalf("expected non-zero code for missing runtime path, got success")
|
t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
payload, ok := out.Data.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected object payload, got %T", out.Data)
|
||||||
|
}
|
||||||
|
results, ok := payload["results"].([]interface{})
|
||||||
|
if !ok || len(results) == 0 {
|
||||||
|
t.Fatalf("expected non-empty results, got %v", payload["results"])
|
||||||
|
}
|
||||||
|
first, ok := results[0].(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected result object, got %T", results[0])
|
||||||
|
}
|
||||||
|
if _, ok := first["message"]; !ok {
|
||||||
|
t.Fatalf("expected message field in diagnosis result")
|
||||||
|
}
|
||||||
|
if got := int(first["fromChainType"].(float64)); got != 1 {
|
||||||
|
t.Fatalf("expected fromChainType=1, got %d", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("tunnel diagnose returns structured payload", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/diagnose", bytes.NewBufferString(`{"tunnelId":`+jsonNumber(tunnelID)+`}`))
|
||||||
|
req.Header.Set("Authorization", adminToken)
|
||||||
|
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 code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
payload, ok := out.Data.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected object payload, got %T", out.Data)
|
||||||
|
}
|
||||||
|
results, ok := payload["results"].([]interface{})
|
||||||
|
if !ok || len(results) == 0 {
|
||||||
|
t.Fatalf("expected non-empty results, got %v", payload["results"])
|
||||||
|
}
|
||||||
|
first, ok := results[0].(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected result object, got %T", results[0])
|
||||||
|
}
|
||||||
|
if _, ok := first["message"]; !ok {
|
||||||
|
t.Fatalf("expected message field in tunnel diagnosis result")
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,138 @@
|
|||||||
|
package contract
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"go-backend/internal/auth"
|
||||||
|
"go-backend/internal/http/response"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestUserTunnelVisibleListContracts(t *testing.T) {
|
||||||
|
secret := "contract-jwt-secret"
|
||||||
|
router, repo := setupDiagnosisContractRouter(t, secret)
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||||
|
VALUES(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||||
|
`, now, now); err != nil {
|
||||||
|
t.Fatalf("insert user: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
insertTunnel := func(name string, status int, inx int64) int64 {
|
||||||
|
res, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
`, name, 1.0, 1, "tls", 99999, now, now, status, nil, inx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("insert tunnel %s: %v", name, err)
|
||||||
|
}
|
||||||
|
id, err := res.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get tunnel id %s: %v", name, err)
|
||||||
|
}
|
||||||
|
return id
|
||||||
|
}
|
||||||
|
|
||||||
|
enabledA := insertTunnel("enabled-A", 1, 1)
|
||||||
|
enabledB := insertTunnel("enabled-B", 1, 2)
|
||||||
|
disabledC := insertTunnel("disabled-C", 0, 3)
|
||||||
|
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||||
|
VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, ?)
|
||||||
|
`, 2, enabledA, 100, 1000, 1, 2727251700000, 0); err != nil {
|
||||||
|
t.Fatalf("insert user_tunnel enabledA: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||||
|
VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, ?)
|
||||||
|
`, 2, enabledB, 100, 1000, 1, 2727251700000, 1); err != nil {
|
||||||
|
t.Fatalf("insert user_tunnel enabledB: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||||
|
VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, ?)
|
||||||
|
`, 2, disabledC, 100, 1000, 1, 2727251700000, 1); err != nil {
|
||||||
|
t.Fatalf("insert user_tunnel disabledC: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate admin token: %v", err)
|
||||||
|
}
|
||||||
|
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate user token: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("admin sees all enabled tunnels without user_tunnel rows", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/tunnel", nil)
|
||||||
|
req.Header.Set("Authorization", adminToken)
|
||||||
|
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 code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
ids := collectTunnelIDs(t, out.Data)
|
||||||
|
if !ids[enabledA] || !ids[enabledB] {
|
||||||
|
t.Fatalf("expected enabled tunnels for admin, got %v", ids)
|
||||||
|
}
|
||||||
|
if ids[disabledC] {
|
||||||
|
t.Fatalf("did not expect disabled tunnel for admin")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("normal user sees enabled assigned tunnels regardless of user_tunnel status", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/tunnel", nil)
|
||||||
|
req.Header.Set("Authorization", userToken)
|
||||||
|
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 code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
ids := collectTunnelIDs(t, out.Data)
|
||||||
|
if !ids[enabledA] || !ids[enabledB] {
|
||||||
|
t.Fatalf("expected enabled assigned tunnels for user, got %v", ids)
|
||||||
|
}
|
||||||
|
if ids[disabledC] {
|
||||||
|
t.Fatalf("did not expect disabled tunnel for user")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func collectTunnelIDs(t *testing.T, data interface{}) map[int64]bool {
|
||||||
|
t.Helper()
|
||||||
|
arr, ok := data.([]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected array data, got %T", data)
|
||||||
|
}
|
||||||
|
ids := make(map[int64]bool, len(arr))
|
||||||
|
for _, item := range arr {
|
||||||
|
obj, ok := item.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected object item, got %T", item)
|
||||||
|
}
|
||||||
|
id := int64(obj["id"].(float64))
|
||||||
|
ids[id] = true
|
||||||
|
}
|
||||||
|
return ids
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user