fix(backend): sync forward ports on tunnel entry change

This commit is contained in:
sagitchu
2026-03-11 11:18:52 +08:00
parent 673d38a089
commit d2a425d761
2 changed files with 121 additions and 0 deletions
@@ -684,6 +684,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
return
}
oldEntryNodeIDs, _ := h.tunnelEntryNodeIDs(id)
h.cleanupTunnelRuntime(id)
h.cleanupFederationRuntime(id)
@@ -755,6 +756,11 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
return
}
newEntryNodeIDs, _ := h.tunnelEntryNodeIDs(id)
if !sameInt64Set(oldEntryNodeIDs, newEntryNodeIDs) {
h.syncTunnelForwardsEntryPorts(id, newEntryNodeIDs)
}
if typeVal == 2 {
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
if applyErr != nil {
@@ -779,6 +785,102 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.OKEmpty())
}
func sameInt64Set(a, b []int64) bool {
if len(a) != len(b) {
return false
}
if len(a) == 0 {
return true
}
m := make(map[int64]int, len(a))
for _, v := range a {
m[v]++
}
for _, v := range b {
c, ok := m[v]
if !ok || c <= 0 {
return false
}
if c == 1 {
delete(m, v)
continue
}
m[v] = c - 1
}
return len(m) == 0
}
func pickForwardPortFromRecords(ports []forwardPortRecord) int {
min := 0
for _, fp := range ports {
if fp.Port <= 0 {
continue
}
if min == 0 || fp.Port < min {
min = fp.Port
}
}
return min
}
func uniqueInt64s(input []int64) []int64 {
if len(input) <= 1 {
return input
}
seen := make(map[int64]struct{}, len(input))
out := make([]int64, 0, len(input))
for _, v := range input {
if _, ok := seen[v]; ok {
continue
}
seen[v] = struct{}{}
out = append(out, v)
}
return out
}
func (h *Handler) syncTunnelForwardsEntryPorts(tunnelID int64, entryNodeIDs []int64) {
if h == nil || h.repo == nil || tunnelID <= 0 {
return
}
entryNodeIDs = uniqueInt64s(entryNodeIDs)
if len(entryNodeIDs) == 0 {
return
}
forwards, err := h.listForwardsByTunnel(tunnelID)
if err != nil || len(forwards) == 0 {
return
}
allowInIP := len(entryNodeIDs) == 1
for i := range forwards {
f := &forwards[i]
if f == nil {
continue
}
oldPorts, err := h.listForwardPorts(f.ID)
if err != nil {
continue
}
port := pickForwardPortFromRecords(oldPorts)
if port <= 0 {
continue
}
var entries []forwardPortReplaceEntry
if allowInIP {
entries = buildForwardPortEntriesWithPreservedInIP(entryNodeIDs, oldPorts, port)
} else {
entries = make([]forwardPortReplaceEntry, 0, len(entryNodeIDs))
for _, nid := range entryNodeIDs {
entries = append(entries, forwardPortReplaceEntry{NodeID: nid, Port: port, InIP: ""})
}
}
_ = h.repo.ReplaceForwardPorts(f.ID, entries)
}
}
func (h *Handler) tunnelDelete(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))