mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-29 07:56:37 +08:00
feat(backend): auto redeploy tunnel and forward config after node upgrade (#180)
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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}))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user