From 0c7b7deaf5502cecadcd0b3b0e2cf04ce25512f7 Mon Sep 17 00:00:00 2001 From: sagit Date: Mon, 9 Feb 2026 05:17:46 +0000 Subject: [PATCH] fix(backend): fix tunnel batch redeploy logic for type 2 tunnels --- .../internal/http/handler/control_plane.go | 10 +- go-backend/internal/http/handler/mutations.go | 96 +++++++++++++++++++ 2 files changed, 104 insertions(+), 2 deletions(-) diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index c06fd02..e7f81df 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -59,6 +59,8 @@ type chainNodeRecord struct { NodeID int64 Port int NodeName string + Protocol string + Strategy string } type diagnosisTarget struct { @@ -862,7 +864,7 @@ func firstPortFromRange(portRange string) int { 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, COALESCE(ct.port, 0), n.name + SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy FROM chain_tunnel ct LEFT JOIN node n ON n.id = ct.node_id WHERE ct.tunnel_id = ? @@ -877,7 +879,9 @@ 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, &item.Port, &name); err != nil { + var protocol sql.NullString + var strategy sql.NullString + if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &item.Port, &name, &protocol, &strategy); err != nil { return nil, err } if strings.TrimSpace(name.String) == "" { @@ -885,6 +889,8 @@ func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, er } else { item.NodeName = name.String } + item.Protocol = defaultString(protocol.String, "tls") + item.Strategy = defaultString(strategy.String, "round") result = append(result, item) } if err := rows.Err(); err != nil { diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index ab4d701..31da320 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -731,6 +731,82 @@ func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail})) } +func (h *Handler) reconstructTunnelState(tunnelID int64) (*tunnelCreateState, error) { + tunnel, err := h.getTunnelRecord(tunnelID) + if err != nil { + return nil, err + } + + chainRows, err := h.listChainNodesForTunnel(tunnelID) + if err != nil { + return nil, err + } + + state := &tunnelCreateState{ + TunnelID: tunnelID, + Type: tunnel.Type, + InNodes: make([]tunnelRuntimeNode, 0), + ChainHops: make([][]tunnelRuntimeNode, 0), + OutNodes: make([]tunnelRuntimeNode, 0), + Nodes: make(map[int64]*nodeRecord), + NodeIDList: make([]int64, 0), + } + + inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows) + + for _, r := range inNodes { + state.InNodes = append(state.InNodes, tunnelRuntimeNode{ + NodeID: r.NodeID, + Protocol: r.Protocol, + Strategy: r.Strategy, + ChainType: 1, + }) + state.NodeIDList = append(state.NodeIDList, r.NodeID) + } + + for _, r := range outNodes { + state.OutNodes = append(state.OutNodes, tunnelRuntimeNode{ + NodeID: r.NodeID, + Protocol: r.Protocol, + Strategy: r.Strategy, + ChainType: 3, + Port: r.Port, + }) + state.NodeIDList = append(state.NodeIDList, r.NodeID) + } + + for _, hop := range chainHops { + stateHop := make([]tunnelRuntimeNode, 0) + for _, r := range hop { + stateHop = append(stateHop, tunnelRuntimeNode{ + NodeID: r.NodeID, + Protocol: r.Protocol, + Strategy: r.Strategy, + ChainType: 2, + Inx: int(r.Inx), + Port: r.Port, + }) + state.NodeIDList = append(state.NodeIDList, r.NodeID) + } + state.ChainHops = append(state.ChainHops, stateHop) + } + + seen := make(map[int64]struct{}) + for _, id := range state.NodeIDList { + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + node, err := h.getNodeRecord(id) + if err != nil { + return nil, err + } + state.Nodes[id] = node + } + + return state, nil +} + func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) { ids := idsFromBody(r, w) if ids == nil { @@ -739,6 +815,26 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) { success := 0 fail := 0 for _, tunnelID := range ids { + tunnel, err := h.getTunnelRecord(tunnelID) + if err != nil { + fail++ + continue + } + + if tunnel.Type == 2 { + h.cleanupTunnelRuntime(tunnelID) + state, err := h.reconstructTunnelState(tunnelID) + if err != nil { + fail++ + continue + } + _, _, applyErr := h.applyTunnelRuntime(state) + if applyErr != nil { + fail++ + continue + } + } + forwards, err := h.listForwardsByTunnel(tunnelID) if err != nil { fail++