fix: protect shared rules across delivery and recovery (#560)

This commit is contained in:
sagit
2026-09-29 14:05:49 +08:00
committed by GitHub
parent 129fa0aa4c
commit c952d2fb3a
49 changed files with 5053 additions and 643 deletions
+302 -119
View File
@@ -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) {