fix(backend): release old port listeners on tunnel switch

Delete stale forward services on old/kept entry nodes during tunnel changes so ports are freed and rebinds don't hit address-in-use.
This commit is contained in:
sagitchu
2026-03-12 10:55:53 +08:00
parent 5e96a8de72
commit 30d9552207
+105 -1
View File
@@ -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}))