mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-11 03:36:37 +08:00
fix: protect shared rules across delivery and recovery (#560)
This commit is contained in:
@@ -827,6 +827,10 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := validateTunnelRuntimeBeforeApply(runtimeState); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if len(runtimeState.InNodes) > 0 {
|
||||
firstNodeID := runtimeState.InNodes[0].NodeID
|
||||
@@ -909,22 +913,27 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
var federationReleaseRefs []federationRuntimeReleaseRef
|
||||
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState, localDomain)
|
||||
if err != nil {
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
applyTunnelPortsToRequest(req, runtimeState)
|
||||
if err := h.replaceTunnelChainsTx(tx, tunnelID, req); err != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.ReplaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); err != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -932,8 +941,11 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
|
||||
if applyErr != nil {
|
||||
h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID, tunnelProtocol)
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
_ = h.deleteTunnelByID(tunnelID)
|
||||
if cleanupErr := h.cleanupFederationRuntime(tunnelID); cleanupErr != nil {
|
||||
applyErr = errors.Join(applyErr, cleanupErr)
|
||||
} else {
|
||||
_ = h.deleteTunnelByID(tunnelID)
|
||||
}
|
||||
response.WriteJSON(w, response.ErrDefault(applyErr.Error()))
|
||||
return
|
||||
}
|
||||
@@ -1091,22 +1103,38 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
probeTargetPort = probeTarget.Port
|
||||
}
|
||||
}
|
||||
oldEntryNodeIDs, _ := h.tunnelEntryNodeIDs(id)
|
||||
oldTunnel, _ := h.getTunnelRecord(id)
|
||||
oldEntryNodeIDs, err := h.tunnelEntryNodeIDs(id)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
oldTunnel, err := h.getTunnelRecord(id)
|
||||
if err != nil || oldTunnel == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
||||
return
|
||||
}
|
||||
if !probeTargetFieldsPresent && oldTunnel != nil {
|
||||
probeTargetHost = oldTunnel.ProbeTargetHost
|
||||
probeTargetPort = oldTunnel.ProbeTargetPort
|
||||
}
|
||||
oldChainRows, _ := h.listChainNodesForTunnel(id)
|
||||
if oldTunnel != nil && oldTunnel.Type == 2 && typeVal != 2 {
|
||||
h.cleanupTunnelRuntime(id)
|
||||
oldChainRows, err := h.listChainNodesForTunnel(id)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
h.cleanupFederationRuntime(id)
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
localDomain := h.federationLocalDomain()
|
||||
|
||||
runtimeState, err := h.prepareTunnelCreateState(h.repo.DB(), req, typeVal, id)
|
||||
// Check the complete request and known database constraints before changing
|
||||
// any existing listener or remote binding.
|
||||
validationTx := h.repo.BeginTx()
|
||||
if validationTx.Error != nil {
|
||||
response.WriteJSON(w, response.Err(-2, validationTx.Error.Error()))
|
||||
return
|
||||
}
|
||||
defer validationTx.Rollback()
|
||||
runtimeState, err := h.prepareTunnelCreateState(validationTx, req, typeVal, id)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
@@ -1119,36 +1147,69 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
|
||||
}
|
||||
}
|
||||
if err := h.validateNftablesTunnelState(entryNodeIDs); err != nil {
|
||||
if err := h.validateNftablesTunnelStateTx(validationTx, entryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.validateTunnelEntryPortConflictsForNewEntriesTx(validationTx, id, oldEntryNodeIDs, entryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := validateTunnelRuntimeBeforeApply(runtimeState); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
|
||||
|
||||
var federationBindings []repo.FederationTunnelBinding
|
||||
var federationReleaseRefs []federationRuntimeReleaseRef
|
||||
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState, localDomain)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
applyTunnelPortsToRequest(req, runtimeState)
|
||||
|
||||
tx := h.repo.BeginTx()
|
||||
if tx.Error != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
response.WriteJSON(w, response.Err(-2, tx.Error.Error()))
|
||||
return
|
||||
}
|
||||
defer func() { tx.Rollback() }()
|
||||
|
||||
updateProtocol := "tls"
|
||||
if len(runtimeState.OutNodes) > 0 && strings.TrimSpace(runtimeState.OutNodes[0].Protocol) != "" {
|
||||
updateProtocol = strings.TrimSpace(runtimeState.OutNodes[0].Protocol)
|
||||
} else if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" {
|
||||
updateProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
|
||||
}
|
||||
if err := h.repo.UpdateTunnelTx(validationTx, id, asString(req["name"]), typeVal,
|
||||
asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1),
|
||||
inIp, ipPreference, updateProtocol, probeTargetHost, probeTargetPort, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.validateTunnelChainReplacementTx(validationTx, runtimeState); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := validationTx.Rollback().Error; err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.cleanupFederationRuntime(id); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if oldTunnel.Type == 2 && typeVal != 2 {
|
||||
h.cleanupTunnelRuntime(id)
|
||||
}
|
||||
|
||||
var federationBindings []repo.FederationTunnelBinding
|
||||
var federationReleaseRefs []federationRuntimeReleaseRef
|
||||
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState, localDomain)
|
||||
if err != nil {
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
applyTunnelPortsToRequest(req, runtimeState)
|
||||
|
||||
tx := h.repo.BeginTx()
|
||||
if tx.Error != nil {
|
||||
tx.Rollback()
|
||||
err := errors.Join(tx.Error, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
defer func() { tx.Rollback() }()
|
||||
|
||||
if err := h.repo.UpdateTunnelTx(
|
||||
tx,
|
||||
id,
|
||||
@@ -1164,21 +1225,27 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
probeTargetPort,
|
||||
now,
|
||||
); err != nil {
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.DeleteChainTunnelsByTunnelTx(tx, id); err != nil {
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.replaceTunnelChainsTx(tx, id, req); err != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.ReplaceFederationTunnelBindingsTx(tx, id, federationBindings); err != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -1190,12 +1257,15 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
if err := h.validateTunnelEntryPortConflictsForNewEntriesTx(tx, id, oldEntryNodeIDs, newEntryNodeIDs); err != nil {
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
tx.Rollback()
|
||||
err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -1220,8 +1290,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if oldTunnel == nil || oldTunnel.Type != 2 {
|
||||
h.rollbackTunnelRuntime(createdChains, createdServices, id, updateProtocol)
|
||||
}
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(id)
|
||||
applyErr = errors.Join(applyErr, h.cleanupFederationRuntime(id))
|
||||
if len(federationReleaseRefs) == 0 && shouldDeferTunnelRuntimeApplyError(applyErr) {
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
return
|
||||
@@ -1430,7 +1499,10 @@ func (h *Handler) validateTunnelEntryPortConflictsForNewEntriesTx(tx *gorm.DB, t
|
||||
}
|
||||
|
||||
forwards, err := h.repo.ListForwardsByTunnelTx(tx, tunnelID)
|
||||
if err != nil || len(forwards) == 0 {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(forwards) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1441,7 +1513,7 @@ func (h *Handler) validateTunnelEntryPortConflictsForNewEntriesTx(tx *gorm.DB, t
|
||||
}
|
||||
oldPorts, portsErr := h.repo.ListForwardPortsTx(tx, f.ID)
|
||||
if portsErr != nil {
|
||||
continue
|
||||
return portsErr
|
||||
}
|
||||
port := pickForwardPortFromRecords(oldPorts)
|
||||
if port <= 0 {
|
||||
@@ -1451,7 +1523,7 @@ func (h *Handler) validateTunnelEntryPortConflictsForNewEntriesTx(tx *gorm.DB, t
|
||||
for _, nodeID := range addedNodeIDs {
|
||||
node, nodeErr := h.repo.GetNodeRecordTx(tx, nodeID)
|
||||
if nodeErr != nil {
|
||||
continue
|
||||
return nodeErr
|
||||
}
|
||||
|
||||
if err := h.validateForwardPortAvailabilityTx(tx, node, port, f.ID); err != nil {
|
||||
@@ -1603,8 +1675,11 @@ func (h *Handler) tunnelDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
if err := h.cleanupFederationRuntime(id); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
h.cleanupTunnelRuntime(id)
|
||||
h.cleanupFederationRuntime(id)
|
||||
if err := h.deleteTunnelByID(id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1671,8 +1746,12 @@ func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, id, tunnelName, err)
|
||||
continue
|
||||
}
|
||||
if err := h.cleanupFederationRuntime(id); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, tunnelName, err)
|
||||
continue
|
||||
}
|
||||
h.cleanupTunnelRuntime(id)
|
||||
h.cleanupFederationRuntime(id)
|
||||
if err := h.deleteTunnelByID(id); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, tunnelName, err)
|
||||
@@ -1771,34 +1850,35 @@ func (h *Handler) redeployTunnelAndForwards(tunnelID int64) error {
|
||||
}
|
||||
|
||||
if tunnel.Type == 2 {
|
||||
h.cleanupTunnelRuntime(tunnelID)
|
||||
h.cleanupFederationRuntime(tunnelID)
|
||||
state, err := h.reconstructTunnelState(tunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateTunnelRuntimeBeforeApply(state); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := h.cleanupFederationRuntime(tunnelID); err != nil {
|
||||
return err
|
||||
}
|
||||
h.cleanupTunnelRuntime(tunnelID)
|
||||
federationBindings, federationReleaseRefs, fedErr := h.applyFederationRuntime(state, h.federationLocalDomain())
|
||||
if fedErr != nil {
|
||||
return fedErr
|
||||
return errors.Join(fedErr, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
}
|
||||
tx := h.repo.BeginTx()
|
||||
if tx.Error != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
return tx.Error
|
||||
return errors.Join(tx.Error, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
}
|
||||
if replaceErr := h.repo.ReplaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); replaceErr != nil {
|
||||
tx.Rollback()
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
return replaceErr
|
||||
return errors.Join(replaceErr, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
}
|
||||
if commitErr := tx.Commit().Error; commitErr != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
return commitErr
|
||||
return errors.Join(commitErr, h.releaseFederationRuntimeRefs(federationReleaseRefs))
|
||||
}
|
||||
_, _, applyErr := h.applyTunnelRuntime(state)
|
||||
if applyErr != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID)
|
||||
applyErr = errors.Join(applyErr, h.cleanupFederationRuntime(tunnelID))
|
||||
return applyErr
|
||||
}
|
||||
}
|
||||
@@ -3393,12 +3473,13 @@ func (h *Handler) prepareTunnelCreateState(tx *gorm.DB, req map[string]interface
|
||||
existingNodeIDs := make(map[int64]struct{})
|
||||
if excludeTunnelID > 0 {
|
||||
var existIDs []int64
|
||||
if err := tx.Model(&model.ChainTunnel{}).
|
||||
Where("tunnel_id = ?", excludeTunnelID).
|
||||
Pluck("node_id", &existIDs).Error; err == nil {
|
||||
for _, eid := range existIDs {
|
||||
existingNodeIDs[eid] = struct{}{}
|
||||
}
|
||||
var err error
|
||||
existIDs, err = h.repo.ListTunnelChainNodeIDsTx(tx, excludeTunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, eid := range existIDs {
|
||||
existingNodeIDs[eid] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3446,6 +3527,59 @@ func (h *Handler) prepareTunnelCreateState(tx *gorm.DB, req map[string]interface
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func validateTunnelRuntimeBeforeApply(state *tunnelCreateState) error {
|
||||
if state.Type != 2 {
|
||||
return nil
|
||||
}
|
||||
for _, node := range state.Nodes {
|
||||
if node.IsRemote == 1 && (strings.TrimSpace(node.RemoteURL) == "" || strings.TrimSpace(node.RemoteToken) == "") {
|
||||
return fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node))
|
||||
}
|
||||
}
|
||||
groups := append([][]tunnelRuntimeNode{state.InNodes}, state.ChainHops...)
|
||||
groups = append(groups, state.OutNodes)
|
||||
for i := 0; i+1 < len(groups); i++ {
|
||||
for _, source := range groups[i] {
|
||||
for _, target := range groups[i+1] {
|
||||
node := state.Nodes[target.NodeID]
|
||||
if node == nil {
|
||||
return errors.New("节点不存在")
|
||||
}
|
||||
if _, err := selectTunnelDialHost(state.Nodes[source.NodeID], node, state.IPPreference, target.ConnectIP); err != nil {
|
||||
return err
|
||||
}
|
||||
if node.IsRemote != 1 && target.Port <= 0 {
|
||||
return errors.New("节点端口不能为空")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Exercise known chain write constraints inside the validation transaction.
|
||||
// Remote ports that have not been reserved yet remain NULL until actual apply.
|
||||
func (h *Handler) validateTunnelChainReplacementTx(tx *gorm.DB, state *tunnelCreateState) error {
|
||||
if err := h.repo.DeleteChainTunnelsByTunnelTx(tx, state.TunnelID); err != nil {
|
||||
return err
|
||||
}
|
||||
groups := append([][]tunnelRuntimeNode{state.InNodes, state.OutNodes}, state.ChainHops...)
|
||||
for groupIndex, group := range groups {
|
||||
for nodeIndex, node := range group {
|
||||
inx := nodeIndex + 1
|
||||
if node.ChainType == 2 {
|
||||
inx = groupIndex - 1
|
||||
}
|
||||
port := sql.NullInt64{Int64: int64(node.Port), Valid: node.Port > 0}
|
||||
if err := h.repo.CreateChainTunnelTx(tx, state.TunnelID, strconv.Itoa(node.ChainType),
|
||||
node.NodeID, port, node.Strategy, inx, node.Protocol, node.ConnectIP); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildTunnelInIP(inNodes []tunnelRuntimeNode, nodes map[int64]*nodeRecord, ipPreference string) string {
|
||||
set := make(map[string]struct{})
|
||||
ordered := make([]string, 0)
|
||||
@@ -3595,8 +3729,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
remoteURL := strings.TrimSpace(node.RemoteURL)
|
||||
remoteToken := strings.TrimSpace(node.RemoteToken)
|
||||
if remoteURL == "" || remoteToken == "" {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node))
|
||||
return bindings, releaseRefs, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node))
|
||||
}
|
||||
|
||||
resourceKey := federationRuntimeResourceKey(state.TunnelID, outNode.NodeID, 3, 0)
|
||||
@@ -3611,10 +3744,14 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
reserveRes, err = fc.ReservePort(remoteURL, remoteToken, localDomain, reserveReq)
|
||||
}
|
||||
if err != nil {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err)
|
||||
return bindings, releaseRefs, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err)
|
||||
}
|
||||
|
||||
releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{
|
||||
RemoteURL: remoteURL, RemoteToken: remoteToken, BindingID: reserveRes.BindingID,
|
||||
ReservationID: reserveRes.ReservationID, ResourceKey: resourceKey,
|
||||
})
|
||||
|
||||
state.OutNodes[outIdx].Port = reserveRes.AllocatedPort
|
||||
outNode = state.OutNodes[outIdx]
|
||||
|
||||
@@ -3627,8 +3764,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
}
|
||||
applyRes, err := fc.ApplyRole(remoteURL, remoteToken, localDomain, applyReq)
|
||||
if err != nil {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err)
|
||||
return bindings, releaseRefs, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err)
|
||||
}
|
||||
if applyRes.AllocatedPort > 0 {
|
||||
state.OutNodes[outIdx].Port = applyRes.AllocatedPort
|
||||
@@ -3648,13 +3784,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
})
|
||||
releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{
|
||||
RemoteURL: remoteURL,
|
||||
RemoteToken: remoteToken,
|
||||
BindingID: applyRes.BindingID,
|
||||
ReservationID: reserveRes.ReservationID,
|
||||
ResourceKey: resourceKey,
|
||||
})
|
||||
releaseRefs[len(releaseRefs)-1].BindingID = defaultString(applyRes.BindingID, reserveRes.BindingID)
|
||||
}
|
||||
|
||||
for hopIdx := len(state.ChainHops) - 1; hopIdx >= 0; hopIdx-- {
|
||||
@@ -3667,8 +3797,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
remoteURL := strings.TrimSpace(node.RemoteURL)
|
||||
remoteToken := strings.TrimSpace(node.RemoteToken)
|
||||
if remoteURL == "" || remoteToken == "" {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node))
|
||||
return bindings, releaseRefs, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node))
|
||||
}
|
||||
|
||||
resourceKey := federationRuntimeResourceKey(state.TunnelID, chainNode.NodeID, 2, hopIdx+1)
|
||||
@@ -3683,10 +3812,14 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
reserveRes, err = fc.ReservePort(remoteURL, remoteToken, localDomain, reserveReq)
|
||||
}
|
||||
if err != nil {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err)
|
||||
return bindings, releaseRefs, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err)
|
||||
}
|
||||
|
||||
releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{
|
||||
RemoteURL: remoteURL, RemoteToken: remoteToken, BindingID: reserveRes.BindingID,
|
||||
ReservationID: reserveRes.ReservationID, ResourceKey: resourceKey,
|
||||
})
|
||||
|
||||
state.ChainHops[hopIdx][nodeIdx].Port = reserveRes.AllocatedPort
|
||||
chainNode = state.ChainHops[hopIdx][nodeIdx]
|
||||
|
||||
@@ -3700,17 +3833,14 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
for _, target := range nextTargets {
|
||||
targetNode := state.Nodes[target.NodeID]
|
||||
if targetNode == nil {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, errors.New("节点不存在")
|
||||
return bindings, releaseRefs, errors.New("节点不存在")
|
||||
}
|
||||
host, hostErr := selectTunnelDialHost(node, targetNode, state.IPPreference, target.ConnectIP)
|
||||
if hostErr != nil {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, hostErr
|
||||
return bindings, releaseRefs, hostErr
|
||||
}
|
||||
if target.Port <= 0 {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, errors.New("节点端口不能为空")
|
||||
return bindings, releaseRefs, errors.New("节点端口不能为空")
|
||||
}
|
||||
applyTargets = append(applyTargets, client.RuntimeTarget{
|
||||
Host: host,
|
||||
@@ -3729,8 +3859,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
}
|
||||
applyRes, err := fc.ApplyRole(remoteURL, remoteToken, localDomain, applyReq)
|
||||
if err != nil {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err)
|
||||
return bindings, releaseRefs, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err)
|
||||
}
|
||||
if applyRes.AllocatedPort > 0 {
|
||||
state.ChainHops[hopIdx][nodeIdx].Port = applyRes.AllocatedPort
|
||||
@@ -3750,70 +3879,124 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
})
|
||||
releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{
|
||||
RemoteURL: remoteURL,
|
||||
RemoteToken: remoteToken,
|
||||
BindingID: applyRes.BindingID,
|
||||
ReservationID: reserveRes.ReservationID,
|
||||
ResourceKey: resourceKey,
|
||||
})
|
||||
releaseRefs[len(releaseRefs)-1].BindingID = defaultString(applyRes.BindingID, reserveRes.BindingID)
|
||||
}
|
||||
}
|
||||
|
||||
return bindings, releaseRefs, nil
|
||||
}
|
||||
|
||||
func (h *Handler) releaseFederationRuntimeRefs(refs []federationRuntimeReleaseRef) {
|
||||
// releaseFederationRuntimeRefs is used after the caller has rolled back its
|
||||
// transaction. Failed releases survive process restarts in a separate queue.
|
||||
func (h *Handler) releaseFederationRuntimeRefs(refs []federationRuntimeReleaseRef) error {
|
||||
if h == nil || len(refs) == 0 {
|
||||
return
|
||||
return nil
|
||||
}
|
||||
fc := client.NewFederationClient()
|
||||
localDomain := h.federationLocalDomain()
|
||||
var failures []error
|
||||
for i := len(refs) - 1; i >= 0; i-- {
|
||||
ref := refs[i]
|
||||
if strings.TrimSpace(ref.RemoteURL) == "" || strings.TrimSpace(ref.RemoteToken) == "" {
|
||||
pending := &model.FederationPendingRelease{
|
||||
RemoteURL: ref.RemoteURL, RemoteToken: ref.RemoteToken,
|
||||
BindingID: ref.BindingID, ReservationID: ref.ReservationID, ResourceKey: ref.ResourceKey,
|
||||
}
|
||||
if err := h.repo.SavePendingFederationRelease(pending); err != nil {
|
||||
failures = append(failures, fmt.Errorf("保存共享运行时清理任务失败: %w", err))
|
||||
continue
|
||||
}
|
||||
req := client.RuntimeReleaseRoleRequest{
|
||||
BindingID: ref.BindingID,
|
||||
ReservationID: ref.ReservationID,
|
||||
ResourceKey: ref.ResourceKey,
|
||||
BindingID: ref.BindingID, ReservationID: ref.ReservationID, ResourceKey: ref.ResourceKey,
|
||||
}
|
||||
if err := fc.ReleaseRole(ref.RemoteURL, ref.RemoteToken, localDomain, req); err != nil {
|
||||
failures = append(failures, err)
|
||||
continue
|
||||
}
|
||||
if err := h.repo.DeletePendingFederationRelease(pending.ID); err != nil {
|
||||
failures = append(failures, err)
|
||||
}
|
||||
_ = fc.ReleaseRole(ref.RemoteURL, ref.RemoteToken, localDomain, req)
|
||||
}
|
||||
return errors.Join(failures...)
|
||||
}
|
||||
|
||||
func (h *Handler) cleanupFederationRuntime(tunnelID int64) {
|
||||
func (h *Handler) cleanupFederationRuntime(tunnelID int64) error {
|
||||
if h == nil || tunnelID <= 0 {
|
||||
return
|
||||
return nil
|
||||
}
|
||||
bindings, err := h.repo.ListActiveFederationTunnelBindingsByTunnel(tunnelID)
|
||||
if err != nil || len(bindings) == 0 {
|
||||
return
|
||||
bindings, err := h.repo.ListFederationTunnelBindingsForCleanup(tunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
fc := client.NewFederationClient()
|
||||
localDomain := h.federationLocalDomain()
|
||||
var failures []error
|
||||
for _, b := range bindings {
|
||||
node, nodeErr := h.repo.GetNodeByID(b.NodeID)
|
||||
if nodeErr != nil || node == nil {
|
||||
// Persist intent before sending the request: timeouts can occur after
|
||||
// the peer has already acted, and must remain safely retryable.
|
||||
if err := h.repo.MarkFederationTunnelBindingPendingRelease(b.ID); err != nil {
|
||||
failures = append(failures, err)
|
||||
continue
|
||||
}
|
||||
remoteURL := strings.TrimSpace(node.RemoteURL.String)
|
||||
if remoteURL == "" {
|
||||
remoteURL = strings.TrimSpace(b.RemoteURL)
|
||||
if err := h.releaseFederationTunnelBinding(b); err != nil {
|
||||
failures = append(failures, err)
|
||||
}
|
||||
remoteToken := strings.TrimSpace(node.RemoteToken.String)
|
||||
if remoteURL == "" || remoteToken == "" {
|
||||
continue
|
||||
}
|
||||
req := client.RuntimeReleaseRoleRequest{
|
||||
BindingID: strings.TrimSpace(b.RemoteBindingID),
|
||||
ResourceKey: strings.TrimSpace(b.ResourceKey),
|
||||
}
|
||||
_ = fc.ReleaseRole(remoteURL, remoteToken, localDomain, req)
|
||||
}
|
||||
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID)
|
||||
return errors.Join(failures...)
|
||||
}
|
||||
|
||||
func (h *Handler) releaseFederationTunnelBinding(b repo.FederationTunnelBinding) error {
|
||||
node, err := h.repo.GetNodeByID(b.NodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if node == nil {
|
||||
return fmt.Errorf("共享节点 %d 不存在,保留待清理绑定", b.NodeID)
|
||||
}
|
||||
remoteURL := strings.TrimSpace(b.RemoteURL)
|
||||
if remoteURL == "" {
|
||||
remoteURL = strings.TrimSpace(node.RemoteURL.String)
|
||||
}
|
||||
remoteToken := strings.TrimSpace(node.RemoteToken.String)
|
||||
if remoteURL == "" || remoteToken == "" {
|
||||
return fmt.Errorf("共享节点 %d 缺少连接配置,保留待清理绑定", b.NodeID)
|
||||
}
|
||||
req := client.RuntimeReleaseRoleRequest{
|
||||
BindingID: strings.TrimSpace(b.RemoteBindingID), ResourceKey: strings.TrimSpace(b.ResourceKey),
|
||||
}
|
||||
if err := client.NewFederationClient().ReleaseRole(remoteURL, remoteToken, h.federationLocalDomain(), req); err != nil {
|
||||
return fmt.Errorf("共享节点 %d 清理失败: %w", b.NodeID, err)
|
||||
}
|
||||
return h.repo.DeleteFederationTunnelBinding(b.ID)
|
||||
}
|
||||
|
||||
// retryPendingFederationRuntimeCleanup never touches active bindings.
|
||||
func (h *Handler) retryPendingFederationRuntimeCleanup() error {
|
||||
bindings, err := h.repo.ListPendingFederationTunnelBindings()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var failures []error
|
||||
for _, binding := range bindings {
|
||||
if err := h.releaseFederationTunnelBinding(binding); err != nil {
|
||||
failures = append(failures, err)
|
||||
}
|
||||
}
|
||||
pending, err := h.repo.ListPendingFederationReleases()
|
||||
if err != nil {
|
||||
return errors.Join(append(failures, err)...)
|
||||
}
|
||||
fc := client.NewFederationClient()
|
||||
for _, item := range pending {
|
||||
req := client.RuntimeReleaseRoleRequest{
|
||||
BindingID: item.BindingID, ReservationID: item.ReservationID, ResourceKey: item.ResourceKey,
|
||||
}
|
||||
if err := fc.ReleaseRole(item.RemoteURL, item.RemoteToken, h.federationLocalDomain(), req); err != nil {
|
||||
failures = append(failures, err)
|
||||
continue
|
||||
}
|
||||
if err := h.repo.DeletePendingFederationRelease(item.ID); err != nil {
|
||||
failures = append(failures, err)
|
||||
}
|
||||
}
|
||||
return errors.Join(failures...)
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64, error) {
|
||||
|
||||
Reference in New Issue
Block a user