mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-03 09:06:36 +08:00
fix: reduce runtime sync disruptions (#484)
This commit is contained in:
@@ -484,13 +484,39 @@ func (h *Handler) forwardServiceBaseCandidates(forward *forwardRecord) ([]string
|
||||
}
|
||||
|
||||
func (h *Handler) deleteForwardServiceBasesOnNode(nodeID int64, bases []string) error {
|
||||
return deleteForwardServiceCandidates(bases, func(name string) error {
|
||||
payload := map[string]interface{}{
|
||||
"services": []string{name},
|
||||
names := buildForwardServiceDeleteNames(bases)
|
||||
if len(names) == 0 {
|
||||
return nil
|
||||
}
|
||||
payload := map[string]interface{}{"services": names}
|
||||
_, err := h.sendNodeCommand(nodeID, "DeleteService", payload, false, true)
|
||||
return err
|
||||
}
|
||||
|
||||
func buildForwardServiceDeleteNames(bases []string) []string {
|
||||
names := make([]string, 0, len(bases)*3)
|
||||
seen := make(map[string]struct{}, len(bases)*3)
|
||||
appendName := func(name string) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return
|
||||
}
|
||||
_, err := h.sendNodeCommand(nodeID, "DeleteService", payload, false, false)
|
||||
return err
|
||||
})
|
||||
if _, ok := seen[name]; ok {
|
||||
return
|
||||
}
|
||||
seen[name] = struct{}{}
|
||||
names = append(names, name)
|
||||
}
|
||||
for _, base := range bases {
|
||||
base = strings.TrimSpace(base)
|
||||
if base == "" {
|
||||
continue
|
||||
}
|
||||
appendName(base + "_tcp")
|
||||
appendName(base + "_udp")
|
||||
appendName(base)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
func (h *Handler) controlForwardServices(forward *forwardRecord, commandType string, tolerateNotFound bool) error {
|
||||
@@ -520,6 +546,25 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str
|
||||
candidateTunnelIDs = append(candidateTunnelIDs, userTunnelIDs...)
|
||||
candidateTunnelIDs = append(candidateTunnelIDs, allUserTunnelIDs...)
|
||||
bases := buildForwardServiceBaseCandidates(forward.ID, forward.UserID, userTunnelID, candidateTunnelIDs)
|
||||
if strings.EqualFold(strings.TrimSpace(commandType), "DeleteService") {
|
||||
seen := map[int64]struct{}{}
|
||||
for _, fp := range ports {
|
||||
if _, ok := seen[fp.NodeID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[fp.NodeID] = struct{}{}
|
||||
if err := h.deleteForwardServiceBasesOnNode(fp.NodeID, bases); err != nil {
|
||||
if isNodeOfflineOrTimeoutError(err) {
|
||||
continue
|
||||
}
|
||||
if tolerateNotFound && isNotFoundError(err) {
|
||||
continue
|
||||
}
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
seen := map[int64]struct{}{}
|
||||
healed := false
|
||||
for _, fp := range ports {
|
||||
|
||||
@@ -201,6 +201,52 @@ func TestDeleteForwardServiceCandidatesDeletesAllMatchingVariants(t *testing.T)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceDeleteNamesBatchesAndDeduplicatesVariants(t *testing.T) {
|
||||
bases := []string{"57_7_7", "57_7_0", "57_7_7"}
|
||||
got := buildForwardServiceDeleteNames(bases)
|
||||
want := []string{"57_7_7_tcp", "57_7_7_udp", "57_7_7", "57_7_0_tcp", "57_7_0_udp", "57_7_0"}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("expected %v, got %v", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemovedTunnelRuntimeNodeIDsSeparatesChainAndServiceRoles(t *testing.T) {
|
||||
oldRows := []chainNodeRecord{
|
||||
{NodeID: 1, ChainType: 1},
|
||||
{NodeID: 2, ChainType: 2},
|
||||
{NodeID: 3, ChainType: 3},
|
||||
{NodeID: 5, ChainType: 2},
|
||||
{NodeID: 6, ChainType: 3},
|
||||
}
|
||||
newRows := []chainNodeRecord{
|
||||
{NodeID: 2, ChainType: 3},
|
||||
{NodeID: 3, ChainType: 3},
|
||||
{NodeID: 5, ChainType: 1},
|
||||
}
|
||||
|
||||
removedChains := removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsChain)
|
||||
if want := []int64{1, 2}; !reflect.DeepEqual(removedChains, want) {
|
||||
t.Fatalf("expected removed chains %v, got %v", want, removedChains)
|
||||
}
|
||||
|
||||
removedServices := removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsService)
|
||||
if want := []int64{5, 6}; !reflect.DeepEqual(removedServices, want) {
|
||||
t.Fatalf("expected removed services %v, got %v", want, removedServices)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelForwardRuntimeNeedsSyncOnlyWhenTypeOrEntriesChange(t *testing.T) {
|
||||
if tunnelForwardRuntimeNeedsSync(2, 2, []int64{1, 2}, []int64{2, 1}) {
|
||||
t.Fatalf("same tunnel type and same entry set should not resync forwards")
|
||||
}
|
||||
if !tunnelForwardRuntimeNeedsSync(1, 2, []int64{1}, []int64{1}) {
|
||||
t.Fatalf("type change should resync forwards")
|
||||
}
|
||||
if !tunnelForwardRuntimeNeedsSync(2, 2, []int64{1}, []int64{1, 2}) {
|
||||
t.Fatalf("entry set change should resync forwards")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateForwardPortAvailabilityRejectsOtherForwardOccupancy(t *testing.T) {
|
||||
h := &Handler{repo: nil}
|
||||
node := &nodeRecord{ID: 9, Name: "test-node"}
|
||||
|
||||
@@ -748,21 +748,8 @@ func (h *Handler) cleanupTunnelRuntime(tunnelID int64) {
|
||||
return
|
||||
}
|
||||
|
||||
protocol := strings.TrimSpace(tunnel.Protocol)
|
||||
if protocol == "" {
|
||||
protocol = "tls"
|
||||
}
|
||||
chainName := fmt.Sprintf("chains_%d", tunnelID)
|
||||
serviceNames := []string{
|
||||
fmt.Sprintf("tunnel_%d", tunnelID),
|
||||
fmt.Sprintf("%d_tls", tunnelID),
|
||||
fmt.Sprintf("%d_kcp", tunnelID),
|
||||
fmt.Sprintf("%d_wss", tunnelID),
|
||||
fmt.Sprintf("%d_mtls", tunnelID),
|
||||
fmt.Sprintf("%d_mwss", tunnelID),
|
||||
fmt.Sprintf("%d_tcp", tunnelID),
|
||||
fmt.Sprintf("%d_mtcp", tunnelID),
|
||||
}
|
||||
serviceNames := tunnelRuntimeServiceNames(tunnelID)
|
||||
|
||||
for _, row := range chainRows {
|
||||
if row.ChainType == 1 {
|
||||
@@ -776,6 +763,70 @@ func (h *Handler) cleanupTunnelRuntime(tunnelID int64) {
|
||||
}
|
||||
}
|
||||
|
||||
func tunnelRuntimeServiceNames(tunnelID int64) []string {
|
||||
return []string{
|
||||
fmt.Sprintf("tunnel_%d", tunnelID),
|
||||
fmt.Sprintf("%d_tls", tunnelID),
|
||||
fmt.Sprintf("%d_kcp", tunnelID),
|
||||
fmt.Sprintf("%d_wss", tunnelID),
|
||||
fmt.Sprintf("%d_mtls", tunnelID),
|
||||
fmt.Sprintf("%d_mwss", tunnelID),
|
||||
fmt.Sprintf("%d_tcp", tunnelID),
|
||||
fmt.Sprintf("%d_mtcp", tunnelID),
|
||||
}
|
||||
}
|
||||
|
||||
func tunnelRuntimeNeedsChain(row chainNodeRecord) bool {
|
||||
return row.ChainType == 1 || row.ChainType == 2
|
||||
}
|
||||
|
||||
func tunnelRuntimeNeedsService(row chainNodeRecord) bool {
|
||||
return row.ChainType == 2 || row.ChainType == 3
|
||||
}
|
||||
|
||||
func removedTunnelRuntimeNodeIDs(oldRows, newRows []chainNodeRecord, needsRuntime func(chainNodeRecord) bool) []int64 {
|
||||
if len(oldRows) == 0 || needsRuntime == nil {
|
||||
return nil
|
||||
}
|
||||
newRuntimeNodes := make(map[int64]struct{}, len(newRows))
|
||||
for _, row := range newRows {
|
||||
if row.NodeID <= 0 || !needsRuntime(row) {
|
||||
continue
|
||||
}
|
||||
newRuntimeNodes[row.NodeID] = struct{}{}
|
||||
}
|
||||
seen := make(map[int64]struct{}, len(oldRows))
|
||||
removed := make([]int64, 0)
|
||||
for _, row := range oldRows {
|
||||
if row.NodeID <= 0 || !needsRuntime(row) {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[row.NodeID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[row.NodeID] = struct{}{}
|
||||
if _, stillNeeded := newRuntimeNodes[row.NodeID]; stillNeeded {
|
||||
continue
|
||||
}
|
||||
removed = append(removed, row.NodeID)
|
||||
}
|
||||
return removed
|
||||
}
|
||||
|
||||
func (h *Handler) cleanupObsoleteTunnelRuntime(tunnelID int64, oldRows, newRows []chainNodeRecord) {
|
||||
if h == nil || tunnelID <= 0 || len(oldRows) == 0 {
|
||||
return
|
||||
}
|
||||
chainName := fmt.Sprintf("chains_%d", tunnelID)
|
||||
for _, nodeID := range removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsChain) {
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true)
|
||||
}
|
||||
serviceNames := tunnelRuntimeServiceNames(tunnelID)
|
||||
for _, nodeID := range removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsService) {
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": serviceNames}, false, true)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelGet(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
@@ -819,12 +870,15 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
oldEntryNodeIDs, _ := h.tunnelEntryNodeIDs(id)
|
||||
|
||||
h.cleanupTunnelRuntime(id)
|
||||
typeVal := asInt(req["type"], 1)
|
||||
oldTunnel, _ := h.getTunnelRecord(id)
|
||||
oldChainRows, _ := h.listChainNodesForTunnel(id)
|
||||
if oldTunnel != nil && oldTunnel.Type == 2 && typeVal != 2 {
|
||||
h.cleanupTunnelRuntime(id)
|
||||
}
|
||||
h.cleanupFederationRuntime(id)
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
typeVal := asInt(req["type"], 1)
|
||||
ipPreference := asString(req["ipPreference"])
|
||||
localDomain := h.federationLocalDomain()
|
||||
|
||||
@@ -917,13 +971,19 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
if typeVal == 2 {
|
||||
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
|
||||
applyRuntime := h.applyTunnelRuntime
|
||||
if oldTunnel != nil && oldTunnel.Type == 2 {
|
||||
applyRuntime = h.applyTunnelRuntimeUpsert
|
||||
}
|
||||
createdChains, createdServices, applyErr := applyRuntime(runtimeState)
|
||||
if applyErr != nil {
|
||||
updateProtocol := "tls"
|
||||
if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" {
|
||||
updateProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
|
||||
}
|
||||
h.rollbackTunnelRuntime(createdChains, createdServices, id, updateProtocol)
|
||||
if oldTunnel == nil || oldTunnel.Type != 2 {
|
||||
h.rollbackTunnelRuntime(createdChains, createdServices, id, updateProtocol)
|
||||
}
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(id)
|
||||
if len(federationReleaseRefs) == 0 && shouldDeferTunnelRuntimeApplyError(applyErr) {
|
||||
@@ -933,9 +993,20 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault(applyErr.Error()))
|
||||
return
|
||||
}
|
||||
newChainRows, _ := h.listChainNodesForTunnel(id)
|
||||
h.cleanupObsoleteTunnelRuntime(id, oldChainRows, newChainRows)
|
||||
}
|
||||
|
||||
if forwards, fwdErr := h.listForwardsByTunnel(id); fwdErr == nil {
|
||||
oldType := 0
|
||||
if oldTunnel != nil {
|
||||
oldType = oldTunnel.Type
|
||||
}
|
||||
if tunnelForwardRuntimeNeedsSync(oldType, typeVal, oldEntryNodeIDs, newEntryNodeIDs) {
|
||||
forwards, fwdErr := h.listForwardsByTunnel(id)
|
||||
if fwdErr != nil {
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
return
|
||||
}
|
||||
for i := range forwards {
|
||||
_ = h.syncForwardServices(&forwards[i], "UpdateService", true)
|
||||
}
|
||||
@@ -944,6 +1015,13 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func tunnelForwardRuntimeNeedsSync(oldType, newType int, oldEntryNodeIDs, newEntryNodeIDs []int64) bool {
|
||||
if oldType != newType {
|
||||
return true
|
||||
}
|
||||
return !sameInt64Set(oldEntryNodeIDs, newEntryNodeIDs)
|
||||
}
|
||||
|
||||
func sameInt64Set(a, b []int64) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
@@ -3358,6 +3436,14 @@ func (h *Handler) cleanupFederationRuntime(tunnelID int64) {
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64, error) {
|
||||
return h.applyTunnelRuntimeWithMode(state, false)
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelRuntimeUpsert(state *tunnelCreateState) ([]int64, []int64, error) {
|
||||
return h.applyTunnelRuntimeWithMode(state, true)
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelRuntimeWithMode(state *tunnelCreateState, upsert bool) ([]int64, []int64, error) {
|
||||
if h == nil || state == nil {
|
||||
return nil, nil, errors.New("invalid tunnel runtime state")
|
||||
}
|
||||
@@ -3376,7 +3462,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
if err != nil {
|
||||
return createdChains, createdServices, err
|
||||
}
|
||||
if _, err := h.sendNodeCommand(inNode.NodeID, "AddChains", chainData, true, false); err != nil {
|
||||
if err := h.applyTunnelChainOnNode(inNode.NodeID, chainData, upsert); err != nil {
|
||||
if shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
@@ -3399,7 +3485,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
if err != nil {
|
||||
return createdChains, createdServices, err
|
||||
}
|
||||
if _, err := h.sendNodeCommand(chainNode.NodeID, "AddChains", chainData, true, false); err != nil {
|
||||
if err := h.applyTunnelChainOnNode(chainNode.NodeID, chainData, upsert); err != nil {
|
||||
if shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
@@ -3408,7 +3494,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
createdChains = append(createdChains, chainNode.NodeID)
|
||||
|
||||
serviceData := buildTunnelChainServiceConfig(state.TunnelID, chainNode, state.Nodes[chainNode.NodeID], len(nextTargets))
|
||||
if err := h.addTunnelServiceOnNode(chainNode.NodeID, state.TunnelID, serviceData); err != nil {
|
||||
if err := h.addTunnelServiceOnNodeWithMode(chainNode.NodeID, state.TunnelID, serviceData, upsert); err != nil {
|
||||
if shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
@@ -3424,7 +3510,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
continue
|
||||
}
|
||||
serviceData := buildTunnelChainServiceConfig(state.TunnelID, outNode, state.Nodes[outNode.NodeID], 1)
|
||||
if err := h.addTunnelServiceOnNode(outNode.NodeID, state.TunnelID, serviceData); err != nil {
|
||||
if err := h.addTunnelServiceOnNodeWithMode(outNode.NodeID, state.TunnelID, serviceData, upsert); err != nil {
|
||||
if shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
@@ -3436,6 +3522,27 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
return createdChains, createdServices, nil
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelChainOnNode(nodeID int64, chainData map[string]interface{}, upsert bool) error {
|
||||
if upsert {
|
||||
return h.upsertTunnelChainOnNode(nodeID, chainData)
|
||||
}
|
||||
_, err := h.sendNodeCommand(nodeID, "AddChains", chainData, true, false)
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *Handler) upsertTunnelChainOnNode(nodeID int64, chainData map[string]interface{}) error {
|
||||
if h == nil {
|
||||
return errors.New("invalid tunnel chain context")
|
||||
}
|
||||
chainName := asString(chainData["name"])
|
||||
if strings.TrimSpace(chainName) == "" {
|
||||
return errors.New("转发链名称不能为空")
|
||||
}
|
||||
payload := map[string]interface{}{"chain": chainName, "data": chainData}
|
||||
_, err := h.sendNodeCommand(nodeID, "UpdateChains", payload, true, false)
|
||||
return err
|
||||
}
|
||||
|
||||
func retryTunnelServiceAddWithCleanup(add func() error, cleanup func() error, wait time.Duration) error {
|
||||
if add == nil {
|
||||
return errors.New("invalid tunnel service add callback")
|
||||
@@ -3457,6 +3564,10 @@ func retryTunnelServiceAddWithCleanup(add func() error, cleanup func() error, wa
|
||||
}
|
||||
|
||||
func (h *Handler) addTunnelServiceOnNode(nodeID, tunnelID int64, serviceData []map[string]interface{}) error {
|
||||
return h.addTunnelServiceOnNodeWithMode(nodeID, tunnelID, serviceData, false)
|
||||
}
|
||||
|
||||
func (h *Handler) addTunnelServiceOnNodeWithMode(nodeID, tunnelID int64, serviceData []map[string]interface{}, upsert bool) error {
|
||||
if h == nil {
|
||||
return errors.New("invalid tunnel service context")
|
||||
}
|
||||
@@ -3466,9 +3577,13 @@ func (h *Handler) addTunnelServiceOnNode(nodeID, tunnelID int64, serviceData []m
|
||||
serviceName = strings.TrimSpace(name)
|
||||
}
|
||||
}
|
||||
command := "AddService"
|
||||
if upsert {
|
||||
command = "UpdateService"
|
||||
}
|
||||
return retryTunnelServiceAddWithCleanup(
|
||||
func() error {
|
||||
_, err := h.sendNodeCommand(nodeID, "AddService", serviceData, true, false)
|
||||
_, err := h.sendNodeCommand(nodeID, command, serviceData, true, false)
|
||||
return err
|
||||
},
|
||||
func() error {
|
||||
@@ -3487,16 +3602,7 @@ func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tu
|
||||
protocol = "tls"
|
||||
}
|
||||
seenServices := make(map[int64]struct{})
|
||||
serviceNames := []string{
|
||||
fmt.Sprintf("tunnel_%d", tunnelID),
|
||||
fmt.Sprintf("%d_tls", tunnelID),
|
||||
fmt.Sprintf("%d_kcp", tunnelID),
|
||||
fmt.Sprintf("%d_wss", tunnelID),
|
||||
fmt.Sprintf("%d_mtls", tunnelID),
|
||||
fmt.Sprintf("%d_mwss", tunnelID),
|
||||
fmt.Sprintf("%d_tcp", tunnelID),
|
||||
fmt.Sprintf("%d_mtcp", tunnelID),
|
||||
}
|
||||
serviceNames := tunnelRuntimeServiceNames(tunnelID)
|
||||
for i := len(serviceNodeIDs) - 1; i >= 0; i-- {
|
||||
nodeID := serviceNodeIDs[i]
|
||||
if _, ok := seenServices[nodeID]; ok {
|
||||
|
||||
Reference in New Issue
Block a user