mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-11 03:36: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
|
jobsCancel context.CancelFunc
|
||||||
jobsStarted bool
|
jobsStarted bool
|
||||||
jobsWG sync.WaitGroup
|
jobsWG sync.WaitGroup
|
||||||
|
|
||||||
|
upgradeMu sync.Mutex
|
||||||
|
pendingUpgradeRedeploy map[int64]struct{}
|
||||||
}
|
}
|
||||||
|
|
||||||
type loginRequest struct {
|
type loginRequest struct {
|
||||||
@@ -70,12 +73,15 @@ type flowItem struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func New(repo *repo.Repository, jwtSecret string) *Handler {
|
func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||||
return &Handler{
|
h := &Handler{
|
||||||
repo: repo,
|
repo: repo,
|
||||||
jwtSecret: jwtSecret,
|
jwtSecret: jwtSecret,
|
||||||
wsServer: ws.NewServer(repo, jwtSecret),
|
wsServer: ws.NewServer(repo, jwtSecret),
|
||||||
captchaTokens: make(map[string]int64),
|
captchaTokens: make(map[string]int64),
|
||||||
|
pendingUpgradeRedeploy: make(map[int64]struct{}),
|
||||||
}
|
}
|
||||||
|
h.wsServer.SetNodeOnlineHook(h.onNodeOnline)
|
||||||
|
return h
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) WebSocketHandler() http.Handler {
|
func (h *Handler) WebSocketHandler() http.Handler {
|
||||||
|
|||||||
@@ -918,6 +918,58 @@ func (h *Handler) reconstructTunnelState(tunnelID int64) (*tunnelCreateState, er
|
|||||||
return state, nil
|
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) {
|
func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||||
ids := idsFromBody(r, w)
|
ids := idsFromBody(r, w)
|
||||||
if ids == nil {
|
if ids == nil {
|
||||||
@@ -926,72 +978,11 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
|||||||
success := 0
|
success := 0
|
||||||
fail := 0
|
fail := 0
|
||||||
for _, tunnelID := range ids {
|
for _, tunnelID := range ids {
|
||||||
tunnel, err := h.getTunnelRecord(tunnelID)
|
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||||
if err != nil {
|
|
||||||
fail++
|
fail++
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
success++
|
||||||
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++
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
|
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)))
|
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("升级失败: %v", err)))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
h.markNodePendingUpgradeRedeploy(req.ID)
|
||||||
|
|
||||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||||
"version": version,
|
"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()}
|
results[index] = upgradeResult{ID: nodeID, Success: false, Message: err.Error()}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
h.markNodePendingUpgradeRedeploy(nodeID)
|
||||||
results[index] = upgradeResult{ID: nodeID, Success: true, Message: result.Message}
|
results[index] = upgradeResult{ID: nodeID, Success: true, Message: result.Message}
|
||||||
}(i, id)
|
}(i, id)
|
||||||
}
|
}
|
||||||
@@ -340,3 +342,66 @@ func (h *Handler) nodeRollback(w http.ResponseWriter, r *http.Request) {
|
|||||||
"message": result.Message,
|
"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
|
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) {
|
func (r *Repository) ListForwardPorts(forwardID int64) ([]model.ForwardPortRecord, error) {
|
||||||
if r == nil || r.db == nil {
|
if r == nil || r.db == nil {
|
||||||
return nil, errors.New("repository not initialized")
|
return nil, errors.New("repository not initialized")
|
||||||
|
|||||||
@@ -68,9 +68,10 @@ type CommandResult struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type Server struct {
|
type Server struct {
|
||||||
repo *repo.Repository
|
repo *repo.Repository
|
||||||
jwtSecret string
|
jwtSecret string
|
||||||
upgrader websocket.Upgrader
|
upgrader websocket.Upgrader
|
||||||
|
onNodeOnline func(nodeID int64)
|
||||||
|
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
admins map[*connWrap]struct{}
|
admins map[*connWrap]struct{}
|
||||||
@@ -79,6 +80,15 @@ type Server struct {
|
|||||||
pending map[string]pendingRequest
|
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 {
|
func NewServer(repo *repo.Repository, jwtSecret string) *Server {
|
||||||
return &Server{
|
return &Server{
|
||||||
repo: repo,
|
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.repo.UpdateNodeOnline(nodeID, 1, version, httpVal, tlsVal, socksVal)
|
||||||
s.broadcastStatus(nodeID, 1)
|
s.broadcastStatus(nodeID, 1)
|
||||||
|
|
||||||
|
s.mu.RLock()
|
||||||
|
onlineHook := s.onNodeOnline
|
||||||
|
s.mu.RUnlock()
|
||||||
|
if onlineHook != nil {
|
||||||
|
go onlineHook(nodeID)
|
||||||
|
}
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
close(done)
|
close(done)
|
||||||
needOfflineBroadcast := false
|
needOfflineBroadcast := false
|
||||||
|
|||||||
Reference in New Issue
Block a user