diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index e010fda..9f0fe2b 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -34,6 +34,9 @@ type Handler struct { jobsCancel context.CancelFunc jobsStarted bool jobsWG sync.WaitGroup + + upgradeMu sync.Mutex + pendingUpgradeRedeploy map[int64]struct{} } type loginRequest struct { @@ -70,12 +73,15 @@ type flowItem struct { } func New(repo *repo.Repository, jwtSecret string) *Handler { - return &Handler{ - repo: repo, - jwtSecret: jwtSecret, - wsServer: ws.NewServer(repo, jwtSecret), - captchaTokens: make(map[string]int64), + h := &Handler{ + repo: repo, + jwtSecret: jwtSecret, + wsServer: ws.NewServer(repo, jwtSecret), + captchaTokens: make(map[string]int64), + pendingUpgradeRedeploy: make(map[int64]struct{}), } + h.wsServer.SetNodeOnlineHook(h.onNodeOnline) + return h } func (h *Handler) WebSocketHandler() http.Handler { diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index cfd628b..71d3eca 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -918,6 +918,58 @@ func (h *Handler) reconstructTunnelState(tunnelID int64) (*tunnelCreateState, er return state, nil } +func (h *Handler) redeployTunnelAndForwards(tunnelID int64) error { + tunnel, err := h.getTunnelRecord(tunnelID) + if err != nil { + return err + } + + if tunnel.Type == 2 { + h.cleanupTunnelRuntime(tunnelID) + h.cleanupFederationRuntime(tunnelID) + state, err := h.reconstructTunnelState(tunnelID) + if err != nil { + return err + } + federationBindings, federationReleaseRefs, fedErr := h.applyFederationRuntime(state, h.federationLocalDomain()) + if fedErr != nil { + return fedErr + } + tx := h.repo.BeginTx() + if tx.Error != nil { + h.releaseFederationRuntimeRefs(federationReleaseRefs) + return tx.Error + } + if replaceErr := h.repo.ReplaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); replaceErr != nil { + tx.Rollback() + h.releaseFederationRuntimeRefs(federationReleaseRefs) + return replaceErr + } + if commitErr := tx.Commit().Error; commitErr != nil { + h.releaseFederationRuntimeRefs(federationReleaseRefs) + return commitErr + } + _, _, applyErr := h.applyTunnelRuntime(state) + if applyErr != nil { + h.releaseFederationRuntimeRefs(federationReleaseRefs) + _ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID) + return applyErr + } + } + + forwards, err := h.listForwardsByTunnel(tunnelID) + if err != nil { + return err + } + for i := range forwards { + if err := h.syncForwardServices(&forwards[i], "UpdateService", true); err != nil { + return err + } + } + + return nil +} + func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) { ids := idsFromBody(r, w) if ids == nil { @@ -926,72 +978,11 @@ 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 { + if err := h.redeployTunnelAndForwards(tunnelID); err != nil { fail++ continue } - - if tunnel.Type == 2 { - h.cleanupTunnelRuntime(tunnelID) - h.cleanupFederationRuntime(tunnelID) - state, err := h.reconstructTunnelState(tunnelID) - if err != nil { - fail++ - continue - } - federationBindings, federationReleaseRefs, fedErr := h.applyFederationRuntime(state, h.federationLocalDomain()) - if fedErr != nil { - fail++ - continue - } - tx := h.repo.BeginTx() - if tx.Error != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) - fail++ - continue - } - if replaceErr := h.repo.ReplaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); replaceErr != nil { - tx.Rollback() - h.releaseFederationRuntimeRefs(federationReleaseRefs) - fail++ - continue - } - if commitErr := tx.Commit().Error; commitErr != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) - fail++ - continue - } - _, _, applyErr := h.applyTunnelRuntime(state) - if applyErr != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) - _ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID) - fail++ - continue - } - } - - forwards, err := h.listForwardsByTunnel(tunnelID) - if err != nil { - fail++ - continue - } - if len(forwards) == 0 { - success++ - continue - } - ok := true - for i := range forwards { - if err := h.syncForwardServices(&forwards[i], "UpdateService", true); err != nil { - ok = false - break - } - } - if ok { - success++ - } else { - fail++ - } + success++ } response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail})) } diff --git a/go-backend/internal/http/handler/upgrade.go b/go-backend/internal/http/handler/upgrade.go index 9fa933e..cf2ae17 100644 --- a/go-backend/internal/http/handler/upgrade.go +++ b/go-backend/internal/http/handler/upgrade.go @@ -166,6 +166,7 @@ func (h *Handler) nodeUpgrade(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.Err(-2, fmt.Sprintf("升级失败: %v", err))) return } + h.markNodePendingUpgradeRedeploy(req.ID) response.WriteJSON(w, response.OK(map[string]interface{}{ "version": version, @@ -246,6 +247,7 @@ func (h *Handler) nodeBatchUpgrade(w http.ResponseWriter, r *http.Request) { results[index] = upgradeResult{ID: nodeID, Success: false, Message: err.Error()} return } + h.markNodePendingUpgradeRedeploy(nodeID) results[index] = upgradeResult{ID: nodeID, Success: true, Message: result.Message} }(i, id) } @@ -340,3 +342,66 @@ func (h *Handler) nodeRollback(w http.ResponseWriter, r *http.Request) { "message": result.Message, })) } + +func (h *Handler) markNodePendingUpgradeRedeploy(nodeID int64) { + if h == nil || nodeID <= 0 { + return + } + h.upgradeMu.Lock() + h.pendingUpgradeRedeploy[nodeID] = struct{}{} + h.upgradeMu.Unlock() +} + +func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool { + if h == nil || nodeID <= 0 { + return false + } + h.upgradeMu.Lock() + _, ok := h.pendingUpgradeRedeploy[nodeID] + if ok { + delete(h.pendingUpgradeRedeploy, nodeID) + } + h.upgradeMu.Unlock() + return ok +} + +func (h *Handler) onNodeOnline(nodeID int64) { + if !h.consumeNodePendingUpgradeRedeploy(nodeID) { + return + } + h.redeployNodeRuntimeAfterUpgrade(nodeID) +} + +func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) { + tunnelIDs, err := h.repo.ListActiveTunnelIDsByNode(nodeID) + if err != nil { + fmt.Printf("post-upgrade redeploy: list tunnels for node %d failed: %v\n", nodeID, err) + return + } + forwardIDs, err := h.repo.ListActiveForwardIDsByNode(nodeID) + if err != nil { + fmt.Printf("post-upgrade redeploy: list forwards for node %d failed: %v\n", nodeID, err) + return + } + + tunnelFailed := make(map[int64]struct{}) + for _, tunnelID := range tunnelIDs { + if err := h.redeployTunnelAndForwards(tunnelID); err != nil { + tunnelFailed[tunnelID] = struct{}{} + fmt.Printf("post-upgrade redeploy: tunnel %d failed on node %d: %v\n", tunnelID, nodeID, err) + } + } + + for _, forwardID := range forwardIDs { + forward, getErr := h.getForwardRecord(forwardID) + if getErr != nil || forward == nil { + continue + } + if _, skipped := tunnelFailed[forward.TunnelID]; skipped { + continue + } + if err := h.syncForwardServices(forward, "UpdateService", true); err != nil { + fmt.Printf("post-upgrade redeploy: forward %d failed on node %d: %v\n", forwardID, nodeID, err) + } + } +} diff --git a/go-backend/internal/store/repo/repository_control.go b/go-backend/internal/store/repo/repository_control.go index 5333180..2aed13c 100644 --- a/go-backend/internal/store/repo/repository_control.go +++ b/go-backend/internal/store/repo/repository_control.go @@ -56,6 +56,40 @@ func (r *Repository) ListForwardsByTunnel(tunnelID int64) ([]model.ForwardRecord return rows, nil } +func (r *Repository) ListActiveTunnelIDsByNode(nodeID int64) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.ChainTunnel{}). + Joins("JOIN tunnel ON tunnel.id = chain_tunnel.tunnel_id"). + Where("chain_tunnel.node_id = ? AND tunnel.status = 1", nodeID). + Select("DISTINCT chain_tunnel.tunnel_id"). + Order("chain_tunnel.tunnel_id ASC"). + Pluck("chain_tunnel.tunnel_id", &ids).Error + if err != nil { + return nil, err + } + return ids, nil +} + +func (r *Repository) ListActiveForwardIDsByNode(nodeID int64) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.ForwardPort{}). + Joins("JOIN forward ON forward.id = forward_port.forward_id"). + Where("forward_port.node_id = ? AND forward.status = 1", nodeID). + Select("DISTINCT forward_port.forward_id"). + Order("forward_port.forward_id ASC"). + Pluck("forward_port.forward_id", &ids).Error + if err != nil { + return nil, err + } + return ids, nil +} + func (r *Repository) ListForwardPorts(forwardID int64) ([]model.ForwardPortRecord, error) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") diff --git a/go-backend/internal/ws/server.go b/go-backend/internal/ws/server.go index 63e194c..7633ef3 100644 --- a/go-backend/internal/ws/server.go +++ b/go-backend/internal/ws/server.go @@ -68,9 +68,10 @@ type CommandResult struct { } type Server struct { - repo *repo.Repository - jwtSecret string - upgrader websocket.Upgrader + repo *repo.Repository + jwtSecret string + upgrader websocket.Upgrader + onNodeOnline func(nodeID int64) mu sync.RWMutex admins map[*connWrap]struct{} @@ -79,6 +80,15 @@ type Server struct { pending map[string]pendingRequest } +func (s *Server) SetNodeOnlineHook(fn func(nodeID int64)) { + if s == nil { + return + } + s.mu.Lock() + s.onNodeOnline = fn + s.mu.Unlock() +} + func NewServer(repo *repo.Repository, jwtSecret string) *Server { return &Server{ repo: repo, @@ -183,6 +193,13 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64 _ = s.repo.UpdateNodeOnline(nodeID, 1, version, httpVal, tlsVal, socksVal) s.broadcastStatus(nodeID, 1) + s.mu.RLock() + onlineHook := s.onNodeOnline + s.mu.RUnlock() + if onlineHook != nil { + go onlineHook(nodeID) + } + defer func() { close(done) needOfflineBroadcast := false