diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index f256cd3..d7d7853 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -836,6 +836,55 @@ func pickForwardPortFromRecords(ports []forwardPortRecord) int { return min } +func forwardPortNodeIDs(ports []forwardPortRecord) []int64 { + ids := make([]int64, 0, len(ports)) + for _, fp := range ports { + if fp.NodeID <= 0 { + continue + } + ids = append(ids, fp.NodeID) + } + return uniqueInt64s(ids) +} + +func (h *Handler) deleteForwardServicesOnNodeBatch(forward *forwardRecord, nodeID int64) error { + if h == nil || forward == nil || nodeID <= 0 { + return errors.New("invalid forward service cleanup context") + } + bases, err := h.forwardServiceBaseCandidates(forward) + if err != nil { + return err + } + if len(bases) == 0 { + return nil + } + + names := make([]string, 0, len(bases)*3) + seen := make(map[string]struct{}, len(bases)*3) + appendName := func(name string) { + if strings.TrimSpace(name) == "" { + return + } + if _, ok := seen[name]; ok { + return + } + seen[name] = struct{}{} + names = append(names, name) + } + for _, base := range bases { + appendName(base + "_tcp") + appendName(base + "_udp") + appendName(base) + } + if len(names) == 0 { + return nil + } + + payload := map[string]interface{}{"services": names} + _, err = h.sendNodeCommand(nodeID, "DeleteService", payload, false, true) + return err +} + func uniqueInt64s(input []int64) []int64 { if len(input) <= 1 { return input @@ -1496,6 +1545,17 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("多入口隧道的转发不支持自定义监听IP")) return } + // When switching tunnels, entry nodes / service base may change. We must clean up old + // listeners on nodes that will be removed, otherwise the old ports keep listening. + tunnelChanged := tunnelID != forward.TunnelID + oldNodeIDs := forwardPortNodeIDs(oldPorts) + newNodeIDs := uniqueInt64s(fwdEntryNodes) + var removedNodeIDs []int64 + var keptNodeIDs []int64 + if tunnelChanged { + removedNodeIDs = diffInt64s(oldNodeIDs, newNodeIDs) + keptNodeIDs = diffInt64s(oldNodeIDs, removedNodeIDs) + } for _, nodeID := range fwdEntryNodes { node, nodeErr := h.getNodeRecord(nodeID) if nodeErr != nil { @@ -1537,12 +1597,41 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.Err(-2, err.Error())) return } - warnings, err := h.syncForwardServicesWithWarnings(updatedForward, "UpdateService", true) + warnings := make([]string, 0) + if tunnelChanged && len(keptNodeIDs) > 0 { + for _, nodeID := range keptNodeIDs { + if delErr := h.deleteForwardServicesOnNodeBatch(forward, nodeID); delErr != nil { + nodeLabel := fmt.Sprintf("%d", nodeID) + if n, nErr := h.getNodeRecord(nodeID); nErr == nil && n != nil && strings.TrimSpace(n.Name) != "" { + nodeLabel = strings.TrimSpace(n.Name) + } + warnings = append(warnings, fmt.Sprintf("节点 %s 清理旧转发监听失败: %v", nodeLabel, delErr)) + } + } + time.Sleep(tunnelServiceBindRetryDelay) + } + + syncWarnings, err := h.syncForwardServicesWithWarnings(updatedForward, "UpdateService", true) if err != nil { h.rollbackForwardMutation(forward, oldPorts) response.WriteJSON(w, response.ErrDefault(err.Error())) return } + warnings = append(warnings, syncWarnings...) + + // Best-effort cleanup for old entry nodes after a successful tunnel switch. + // Avoid cleaning nodes that are still used by the updated forward. + if tunnelChanged && len(removedNodeIDs) > 0 { + for _, nodeID := range removedNodeIDs { + if delErr := h.deleteForwardServicesOnNodeBatch(forward, nodeID); delErr != nil { + nodeLabel := fmt.Sprintf("%d", nodeID) + if n, nErr := h.getNodeRecord(nodeID); nErr == nil && n != nil && strings.TrimSpace(n.Name) != "" { + nodeLabel = strings.TrimSpace(n.Name) + } + warnings = append(warnings, fmt.Sprintf("节点 %s 清理旧隧道残留服务失败: %v", nodeLabel, delErr)) + } + } + } if len(warnings) > 0 { response.WriteJSON(w, response.OK(map[string]interface{}{"warnings": warnings})) return @@ -1845,6 +1934,7 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques fail++ continue } + oldNodeIDs := forwardPortNodeIDs(oldPorts) port := h.repo.GetMinForwardPort(id) if err := h.repo.UpdateForwardTunnel(id, req.TargetTunnelID, time.Now().UnixMilli()); err != nil { fail++ @@ -1858,6 +1948,9 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques p = h.pickTunnelPort(req.TargetTunnelID) } bctEntryNodes, _ := h.tunnelEntryNodeIDs(req.TargetTunnelID) + newNodeIDs := uniqueInt64s(bctEntryNodes) + removedNodeIDs := diffInt64s(oldNodeIDs, newNodeIDs) + keptNodeIDs := diffInt64s(oldNodeIDs, removedNodeIDs) portRangeOk := true for _, nid := range bctEntryNodes { nd, ndErr := h.getNodeRecord(nid) @@ -1884,11 +1977,22 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques fail++ continue } + if len(keptNodeIDs) > 0 { + for _, nodeID := range keptNodeIDs { + _ = h.deleteForwardServicesOnNodeBatch(forward, nodeID) + } + time.Sleep(tunnelServiceBindRetryDelay) + } if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil { h.rollbackForwardMutation(forward, oldPorts) fail++ continue } + if len(removedNodeIDs) > 0 { + for _, nodeID := range removedNodeIDs { + _ = h.deleteForwardServicesOnNodeBatch(forward, nodeID) + } + } success++ } response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))