From 1a8f424d5397559f8877e37fa1f0ad9c8b0b0ac8 Mon Sep 17 00:00:00 2001 From: sagit Date: Sat, 7 Feb 2026 12:48:41 +0000 Subject: [PATCH] fix: complete diagnosis parity and tunnel visibility on Go backend --- .../internal/http/handler/control_plane.go | 459 ++++++++++++++---- go-backend/internal/http/handler/handler.go | 9 +- .../internal/store/sqlite/repository.go | 36 +- .../tests/contract/diagnosis_contract_test.go | 239 +++++++++ .../tests/contract/forward_contract_test.go | 80 ++- .../tunnel_visibility_contract_test.go | 138 ++++++ 6 files changed, 854 insertions(+), 107 deletions(-) create mode 100644 go-backend/tests/contract/diagnosis_contract_test.go create mode 100644 go-backend/tests/contract/tunnel_visibility_contract_test.go diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 7bf1f6a..02905d8 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -6,6 +6,7 @@ import ( "fmt" "net" "net/http" + "sort" "strconv" "strings" "time" @@ -43,6 +44,7 @@ type nodeRecord struct { ID int64 Name string ServerIP string + PortRange string TCPListenAddr string UDPListenAddr string InterfaceName string @@ -52,9 +54,16 @@ type chainNodeRecord struct { ChainType int Inx int64 NodeID int64 + Port int NodeName string } +type diagnosisTarget struct { + Address string + IP string + Port int +} + func (h *Handler) resolveForwardAccess(r *http.Request, forwardID int64) (*forwardRecord, int64, int, error) { userID, roleID, err := userRoleFromRequest(r) if err != nil { @@ -183,22 +192,24 @@ 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, tcp_listen_addr, udp_listen_addr, interface_name + SELECT id, name, server_ip, port, tcp_listen_addr, udp_listen_addr, interface_name FROM node WHERE id = ? LIMIT 1 `, nodeID) var n nodeRecord + 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, &tcpListen, &udpListen, &iface) + err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &portRange, &tcpListen, &udpListen, &iface) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, errors.New("节点不存在") } return nil, err } + n.PortRange = strings.TrimSpace(portRange.String) n.TCPListenAddr = strings.TrimSpace(tcpListen.String) n.UDPListenAddr = strings.TrimSpace(udpListen.String) n.InterfaceName = strings.TrimSpace(iface.String) @@ -343,68 +354,102 @@ func (h *Handler) diagnoseForwardRuntime(forward *forwardRecord) (map[string]int if forward == nil { return nil, errForwardNotFound } - targets := splitRemoteTargets(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) + targets, err := resolveDiagnosisTargets(forward.RemoteAddr) if err != nil { 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)) - seen := map[int64]struct{}{} - for _, fp := range ports { - if _, ok := seen[fp.NodeID]; ok { - continue - } - seen[fp.NodeID] = struct{}{} - node, nodeErr := h.getNodeRecord(fp.NodeID) - nodeName := fmt.Sprintf("node_%d", fp.NodeID) - if nodeErr == nil { - nodeName = node.Name - } + chainRows, err := h.listChainNodesForTunnel(forward.TunnelID) + if err != nil { + return nil, err + } + if len(chainRows) == 0 { + return nil, errors.New("隧道配置不完整") + } - resultItem := map[string]interface{}{ - "nodeName": nodeName, - "nodeId": strconv.FormatInt(fp.NodeID, 10), - "targetIp": ip, - "targetPort": port, - "averageTime": 0, - "packetLoss": 100, - } + inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows) + results := make([]map[string]interface{}, 0, len(chainRows)*2+len(targets)) + nodeCache := map[int64]*nodeRecord{} - if nodeErr != nil { - resultItem["success"] = false - resultItem["description"] = "节点信息读取失败" - resultItem["errorMessage"] = nodeErr.Error() - results = append(results, resultItem) - continue - } - - pingData, pingErr := h.tcpPingViaNode(fp.NodeID, ip, port) - if pingErr != nil { - resultItem["success"] = false - resultItem["description"] = "诊断失败" - resultItem["errorMessage"] = pingErr.Error() - } else { - resultItem["success"] = asBool(pingData["success"], false) - resultItem["description"] = "诊断完成" - resultItem["averageTime"] = asFloat(pingData["averageTime"], 0) - resultItem["packetLoss"] = asFloat(pingData["packetLoss"], 100) - if msg := asString(pingData["errorMessage"]); msg != "" { - resultItem["errorMessage"] = msg + switch tunnel.Type { + case 1: + 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, + }) + } + } + case 2: + for _, inNode := range inNodes { + if len(chainHops) > 0 { + for _, firstNode := range chainHops[0] { + description := fmt.Sprintf("入口(%s)->第1跳(%s)", inNode.NodeName, firstNode.NodeName) + h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, firstNode, description, map[string]interface{}{ + "fromChainType": 1, + "toChainType": 2, + "toInx": firstNode.Inx, + }) + } + } 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{}{ @@ -437,53 +482,78 @@ func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{}, return nil, errors.New("隧道配置不完整") } - results := make([]map[string]interface{}, 0, len(chainRows)) - seen := map[int64]struct{}{} - for _, row := range chainRows { - if _, ok := seen[row.NodeID]; ok { - continue - } - seen[row.NodeID] = struct{}{} + inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows) + results := make([]map[string]interface{}, 0, len(chainRows)*2) + nodeCache := map[int64]*nodeRecord{} - desc := fmt.Sprintf("节点(%s)->外网", row.NodeName) - switch row.ChainType { - case 1: - desc = fmt.Sprintf("入口(%s)->外网", row.NodeName) - case 2: - desc = fmt.Sprintf("中继(%s)->外网", row.NodeName) - case 3: - desc = fmt.Sprintf("出口(%s)->外网", row.NodeName) + switch tunnel.Type { + case 1: + 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, + }) } - - resultItem := map[string]interface{}{ - "nodeName": row.NodeName, - "nodeId": strconv.FormatInt(row.NodeID, 10), - "targetIp": "www.google.com", - "targetPort": 443, - "averageTime": 0, - "packetLoss": 100, - "fromChainType": row.ChainType, - "description": desc, - } - if row.ChainType == 2 { - resultItem["fromInx"] = row.Inx - } - - pingData, pingErr := h.tcpPingViaNode(row.NodeID, "www.google.com", 443) - if pingErr != nil { - 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 + case 2: + for _, inNode := range inNodes { + if len(chainHops) > 0 { + for _, firstNode := range chainHops[0] { + description := fmt.Sprintf("入口(%s)->第1跳(%s)", inNode.NodeName, firstNode.NodeName) + h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, firstNode, description, map[string]interface{}{ + "fromChainType": 1, + "toChainType": 2, + "toInx": firstNode.Inx, + }) + } + } 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, + }) + } } } - 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{}{ @@ -495,9 +565,198 @@ func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{}, 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) { 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 LEFT JOIN node n ON n.id = ct.node_id WHERE ct.tunnel_id = ? @@ -512,7 +771,7 @@ func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, er for rows.Next() { var item chainNodeRecord 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 } if strings.TrimSpace(name.String) == "" { diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index b4cc0a0..570e961 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -443,13 +443,18 @@ func (h *Handler) userTunnelVisibleList(w http.ResponseWriter, r *http.Request) return } - userID, err := userIDFromRequest(r) + userID, roleID, err := userRoleFromRequest(r) if err != nil { response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) 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 { response.WriteJSON(w, response.Err(-2, err.Error())) return diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index 55baf77..2c81562 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -656,10 +656,10 @@ func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]inter } rows, err := r.db.Query(` - SELECT t.id, t.name + SELECT DISTINCT t.id, t.name FROM user_tunnel ut 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 `, userID) if err != nil { @@ -683,6 +683,38 @@ func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]inter 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) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") diff --git a/go-backend/tests/contract/diagnosis_contract_test.go b/go-backend/tests/contract/diagnosis_contract_test.go new file mode 100644 index 0000000..cc3c76a --- /dev/null +++ b/go-backend/tests/contract/diagnosis_contract_test.go @@ -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 +} diff --git a/go-backend/tests/contract/forward_contract_test.go b/go-backend/tests/contract/forward_contract_test.go index 0b9973f..7868eb9 100644 --- a/go-backend/tests/contract/forward_contract_test.go +++ b/go-backend/tests/contract/forward_contract_test.go @@ -37,6 +37,25 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) { 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(` 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, ?) @@ -65,6 +84,10 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) { 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("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)+`}`)) @@ -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.Header.Set("Authorization", userToken) res := httptest.NewRecorder() @@ -117,8 +140,59 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) { if err := json.NewDecoder(res.Body).Decode(&out); err != nil { t.Fatalf("decode response: %v", err) } - if out.Code == 0 { - t.Fatalf("expected non-zero code for missing runtime path, got success") + 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 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") } }) } diff --git a/go-backend/tests/contract/tunnel_visibility_contract_test.go b/go-backend/tests/contract/tunnel_visibility_contract_test.go new file mode 100644 index 0000000..c623675 --- /dev/null +++ b/go-backend/tests/contract/tunnel_visibility_contract_test.go @@ -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 +}