fix: preserve bind IP when editing forward without explicit inIp change

- Add replaceForwardPortsPreservingInIP to maintain existing InIP values
- Track inIpTouched state in frontend to distinguish user changes
- Only send inIp in update request when user explicitly changed it
- Add unit tests for forward port replacement with preserved InIP
This commit is contained in:
sagitchu
2026-03-03 20:22:58 +08:00
parent 75edeb9afa
commit 1450b25475
5 changed files with 129 additions and 29 deletions
+58 -22
View File
@@ -1285,7 +1285,12 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
port = h.pickTunnelPort(tunnelID) port = h.pickTunnelPort(tunnelID)
} }
} }
inIp := asString(req["inIp"]) hasInIP := false
inIp := ""
if rawInIP, ok := req["inIp"]; ok {
hasInIP = true
inIp = asString(rawInIP)
}
fwdEntryNodes, _ := h.tunnelEntryNodeIDs(tunnelID) fwdEntryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
for _, nodeID := range fwdEntryNodes { for _, nodeID := range fwdEntryNodes {
node, nodeErr := h.getNodeRecord(nodeID) node, nodeErr := h.getNodeRecord(nodeID)
@@ -1302,7 +1307,14 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
if err := h.replaceForwardPorts(id, tunnelID, port, inIp); err != nil { if hasInIP {
err = h.replaceForwardPorts(id, tunnelID, port, inIp)
} else if tunnelID != forward.TunnelID {
err = h.replaceForwardPorts(id, tunnelID, port, "")
} else {
err = h.replaceForwardPortsPreservingInIP(id, tunnelID, port, oldPorts)
}
if err != nil {
h.rollbackForwardMutation(forward, oldPorts) h.rollbackForwardMutation(forward, oldPorts)
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
@@ -3016,38 +3028,62 @@ func parsePorts(portRange string) ([]int, error) {
return ports, nil return ports, nil
} }
type forwardPortReplaceEntry = struct {
NodeID int64
Port int
InIP string
}
func buildForwardPortEntriesWithPreservedInIP(entryNodeIDs []int64, oldPorts []forwardPortRecord, port int) []forwardPortReplaceEntry {
preservedByNode := make(map[int64]string)
for _, fp := range oldPorts {
current, exists := preservedByNode[fp.NodeID]
if !exists {
preservedByNode[fp.NodeID] = fp.InIP
continue
}
if strings.TrimSpace(current) == "" && strings.TrimSpace(fp.InIP) != "" {
preservedByNode[fp.NodeID] = fp.InIP
}
}
entries := make([]forwardPortReplaceEntry, 0, len(entryNodeIDs))
for _, nid := range entryNodeIDs {
entries = append(entries, forwardPortReplaceEntry{
NodeID: nid,
Port: port,
InIP: preservedByNode[nid],
})
}
return entries
}
func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int, inIp string) error { func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int, inIp string) error {
entryNodes, err := h.tunnelEntryNodeIDs(tunnelID) entryNodes, err := h.tunnelEntryNodeIDs(tunnelID)
if err != nil { if err != nil {
return err return err
} }
entries := make([]struct { entries := make([]forwardPortReplaceEntry, len(entryNodes))
NodeID int64
Port int
InIP string
}, len(entryNodes))
for i, nid := range entryNodes { for i, nid := range entryNodes {
entries[i] = struct { entries[i] = forwardPortReplaceEntry{NodeID: nid, Port: port, InIP: inIp}
NodeID int64
Port int
InIP string
}{NodeID: nid, Port: port, InIP: inIp}
} }
return h.repo.ReplaceForwardPorts(forwardID, entries) return h.repo.ReplaceForwardPorts(forwardID, entries)
} }
func (h *Handler) replaceForwardPortsPreservingInIP(forwardID, tunnelID int64, port int, oldPorts []forwardPortRecord) error {
entryNodes, err := h.tunnelEntryNodeIDs(tunnelID)
if err != nil {
return err
}
entries := buildForwardPortEntriesWithPreservedInIP(entryNodes, oldPorts, port)
return h.repo.ReplaceForwardPorts(forwardID, entries)
}
func (h *Handler) replaceForwardPortsWithRecords(forwardID int64, ports []forwardPortRecord) error { func (h *Handler) replaceForwardPortsWithRecords(forwardID int64, ports []forwardPortRecord) error {
entries := make([]struct { entries := make([]forwardPortReplaceEntry, len(ports))
NodeID int64
Port int
InIP string
}, len(ports))
for i, fp := range ports { for i, fp := range ports {
entries[i] = struct { entries[i] = forwardPortReplaceEntry{NodeID: fp.NodeID, Port: fp.Port, InIP: fp.InIP}
NodeID int64
Port int
InIP string
}{NodeID: fp.NodeID, Port: fp.Port, InIP: fp.InIP}
} }
return h.repo.ReplaceForwardPorts(forwardID, entries) return h.repo.ReplaceForwardPorts(forwardID, entries)
} }
@@ -0,0 +1,39 @@
package handler
import "testing"
func TestBuildForwardPortEntriesWithPreservedInIP(t *testing.T) {
entryNodeIDs := []int64{10, 20, 30}
oldPorts := []forwardPortRecord{
{NodeID: 10, Port: 10001, InIP: ""},
{NodeID: 10, Port: 10002, InIP: "10.0.0.10"},
{NodeID: 20, Port: 10003, InIP: "10.0.0.20"},
}
entries := buildForwardPortEntriesWithPreservedInIP(entryNodeIDs, oldPorts, 18080)
if len(entries) != 3 {
t.Fatalf("expected 3 entries, got %d", len(entries))
}
if entries[0].NodeID != 10 || entries[0].Port != 18080 || entries[0].InIP != "10.0.0.10" {
t.Fatalf("unexpected first entry: %+v", entries[0])
}
if entries[1].NodeID != 20 || entries[1].Port != 18080 || entries[1].InIP != "10.0.0.20" {
t.Fatalf("unexpected second entry: %+v", entries[1])
}
if entries[2].NodeID != 30 || entries[2].Port != 18080 || entries[2].InIP != "" {
t.Fatalf("unexpected third entry: %+v", entries[2])
}
}
func TestBuildForwardPortEntriesWithPreservedInIP_EmptyOldPorts(t *testing.T) {
entryNodeIDs := []int64{99}
entries := buildForwardPortEntriesWithPreservedInIP(entryNodeIDs, nil, 17000)
if len(entries) != 1 {
t.Fatalf("expected 1 entry, got %d", len(entries))
}
if entries[0].NodeID != 99 || entries[0].Port != 17000 || entries[0].InIP != "" {
t.Fatalf("unexpected entry: %+v", entries[0])
}
}
@@ -0,0 +1,11 @@
# 003 Forward Edit Bind IP Preserve
## Checklist
- [x] Confirm forward edit flow and identify why untouched listen IP gets overwritten.
- [x] Update frontend forward edit submit logic to only send `inIp` when user explicitly changes listen IP.
- [x] On tunnel switch in edit form, reset listen IP to default unless user reselects.
- [x] Update backend forward update logic to preserve existing `forward_port.in_ip` when request omits `inIp` and tunnel is unchanged.
- [x] Keep backend behavior explicit: if `inIp` is sent (including empty), apply requested value; if tunnel changed with no `inIp`, use default bind.
- [x] Add regression tests for preserved bind-IP reconstruction helper behavior.
- [x] Run focused frontend/backend checks for touched files.
+17 -6
View File
@@ -586,6 +586,7 @@ export default function ForwardPage() {
strategy: "fifo", strategy: "fifo",
speedId: null, speedId: null,
}); });
const [inIpTouched, setInIpTouched] = useState(false);
// 表单验证错误 // 表单验证错误
const [errors, setErrors] = useState<{ [key: string]: string }>({}); const [errors, setErrors] = useState<{ [key: string]: string }>({});
@@ -1276,6 +1277,7 @@ export default function ForwardPage() {
// 新增转发 // 新增转发
const handleAdd = () => { const handleAdd = () => {
setIsEdit(false); setIsEdit(false);
setInIpTouched(false);
setForm({ setForm({
name: "", name: "",
tunnelId: null, tunnelId: null,
@@ -1293,6 +1295,7 @@ export default function ForwardPage() {
// 编辑转发 // 编辑转发
const handleEdit = (forward: Forward) => { const handleEdit = (forward: Forward) => {
setIsEdit(true); setIsEdit(true);
setInIpTouched(false);
setForm({ setForm({
id: forward.id, id: forward.id,
userId: forward.userId, userId: forward.userId,
@@ -1357,11 +1360,17 @@ export default function ForwardPage() {
const nextTunnelId = parseInt(tunnelId); const nextTunnelId = parseInt(tunnelId);
const options = tunnelInIpOptionMap.get(nextTunnelId) || []; const options = tunnelInIpOptionMap.get(nextTunnelId) || [];
setForm((prev) => ({ setInIpTouched(false);
...prev,
tunnelId: nextTunnelId, setForm((prev) => {
inIp: options.includes(prev.inIp) ? prev.inIp : "", const tunnelChanged = prev.tunnelId !== nextTunnelId;
}));
return {
...prev,
tunnelId: nextTunnelId,
inIp: tunnelChanged ? "" : options.includes(prev.inIp) ? prev.inIp : "",
};
});
}; };
// 提交表单 // 提交表单
@@ -1388,7 +1397,7 @@ export default function ForwardPage() {
name: form.name, name: form.name,
tunnelId: form.tunnelId, tunnelId: form.tunnelId,
inPort: form.inPort, inPort: form.inPort,
inIp: form.inIp || undefined, ...(inIpTouched ? { inIp: form.inIp || "" } : {}),
remoteAddr: processedRemoteAddr, remoteAddr: processedRemoteAddr,
strategy: addressCount > 1 ? form.strategy : "fifo", strategy: addressCount > 1 ? form.strategy : "fifo",
speedId: normalizeSpeedId(form.speedId), speedId: normalizeSpeedId(form.speedId),
@@ -4288,6 +4297,8 @@ export default function ForwardPage() {
onSelectionChange={(keys) => { onSelectionChange={(keys) => {
const selectedKey = Array.from(keys)[0] as string; const selectedKey = Array.from(keys)[0] as string;
setInIpTouched(true);
setForm((prev) => ({ setForm((prev) => ({
...prev, ...prev,
inIp: selectedKey === "__default__" ? "" : selectedKey, inIp: selectedKey === "__default__" ? "" : selectedKey,
@@ -115,7 +115,10 @@ const validateNyItem = (line: string, value: unknown): ParsedNyImportLine => {
return { line, error: "listen_port格式错误,应为1-65535之间的数字" }; return { line, error: "listen_port格式错误,应为1-65535之间的数字" };
} }
if (normalizedListenPort !== null && !isValidListenPort(normalizedListenPort)) { if (
normalizedListenPort !== null &&
!isValidListenPort(normalizedListenPort)
) {
return { line, error: "listen_port必须为1-65535之间的数字" }; return { line, error: "listen_port必须为1-65535之间的数字" };
} }