From b8193417f5988052008ad9ba3f9239c5288a0fe1 Mon Sep 17 00:00:00 2001 From: sagitchu Date: Tue, 3 Mar 2026 08:24:15 +0800 Subject: [PATCH] feat: add custom IP selection for nodes, tunnels, and forwards - Add extra_ips field to nodes for multi-IP servers - Add connect_ip field to chain_tunnel for specifying connection address - Add in_ip field to forward_port for specifying listen address - Frontend: add UI controls for extra IPs on node form - Frontend: add connect IP input for tunnel chain nodes - Frontend: add listen IP input for forward form - Backend: resolve forward ingress with custom listen IP priority Entire-Checkpoint: 557563462c16 --- .../internal/http/handler/control_plane.go | 8 +- .../internal/http/handler/dual_stack_test.go | 73 +++++++++-------- go-backend/internal/http/handler/mutations.go | 37 +++++++-- go-backend/internal/store/model/model.go | 15 +++- go-backend/internal/store/repo/repository.go | 55 +++++++------ .../internal/store/repo/repository_control.go | 15 +++- .../store/repo/repository_mutations.go | 20 +++-- vite-frontend/src/pages/forward.tsx | 18 +++- vite-frontend/src/pages/node.tsx | 19 ++++- vite-frontend/src/pages/tunnel.tsx | 82 +++++++++++++++++++ 10 files changed, 255 insertions(+), 87 deletions(-) diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 7ec8976..5dc4b33 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -72,7 +72,7 @@ func (h *Handler) buildDiagnosisStreamStartItems(workItems []diagnosisWorkItem) fromNode, _ := h.cachedNode(nodeCache, workItem.fromNodeID) targetNode, err := h.cachedNode(nodeCache, workItem.toNode.NodeID) if err == nil { - resolvedIP, resolvedPort, resolveErr := resolveChainProbeTarget(fromNode, targetNode, workItem.toNode.Port, workItem.ipPreference) + resolvedIP, resolvedPort, resolveErr := resolveChainProbeTarget(fromNode, targetNode, workItem.toNode.Port, workItem.ipPreference, "") if resolveErr == nil { targetIP = resolvedIP targetPort = resolvedPort @@ -1099,7 +1099,7 @@ func (h *Handler) appendChainHopDiagnosis(results *[]map[string]interface{}, nod h.appendFailedDiagnosis(results, nodeCache, fromNodeID, "", 0, description, metadata, err.Error()) return } - targetIP, targetPort, err := resolveChainProbeTarget(fromNode, targetNode, toNode.Port, ipPreference) + targetIP, targetPort, err := resolveChainProbeTarget(fromNode, targetNode, toNode.Port, ipPreference, "") if err != nil { h.appendFailedDiagnosis(results, nodeCache, fromNodeID, strings.Trim(strings.TrimSpace(targetNode.ServerIP), "[]"), toNode.Port, description, metadata, err.Error()) return @@ -1107,11 +1107,11 @@ func (h *Handler) appendChainHopDiagnosis(results *[]map[string]interface{}, nod h.appendPathDiagnosis(results, nodeCache, fromNodeID, targetIP, targetPort, description, metadata, options) } -func resolveChainProbeTarget(fromNode, targetNode *nodeRecord, preferredPort int, ipPreference string) (string, int, error) { +func resolveChainProbeTarget(fromNode, targetNode *nodeRecord, preferredPort int, ipPreference string, connectIp string) (string, int, error) { if targetNode == nil { return "", 0, errors.New("目标节点不存在") } - host, err := selectTunnelDialHost(fromNode, targetNode, ipPreference) + host, err := selectTunnelDialHost(fromNode, targetNode, ipPreference, connectIp) if err != nil { host = strings.Trim(strings.TrimSpace(targetNode.ServerIP), "[]") } diff --git a/go-backend/internal/http/handler/dual_stack_test.go b/go-backend/internal/http/handler/dual_stack_test.go index ee50353..3c423a4 100644 --- a/go-backend/internal/http/handler/dual_stack_test.go +++ b/go-backend/internal/http/handler/dual_stack_test.go @@ -8,9 +8,25 @@ import ( // nodeSupportsV4 / nodeSupportsV6 // --------------------------------------------------------------------------- -func TestNodeSupportsV4_Nil(t *testing.T) { - if nodeSupportsV4(nil) { - t.Fatal("nil node must not support v4") +func TestSelectTunnelDialHost_ConnectIpPriority(t *testing.T) { + from := dualStackNode("from", "10.0.0.1", "2001:db8::1") + to := dualStackNode("to", "10.0.0.2", "2001:db8::2") + + // Empty connectIp should be ignored, IP preference takes effect + host, err := selectTunnelDialHost(from, to, "", "") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if host != "10.0.0.2" { + t.Fatalf("empty connectIp should be ignored (v4 preference applies), got %q", host) + } + // Non-empty connectIp should override IP preference + host, err = selectTunnelDialHost(from, to, "v6", "192.168.0.3") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if host != "192.168.0.3" { + t.Fatalf("connectIp should override v6 preference, got %q", host) } } @@ -23,14 +39,14 @@ func TestNodeSupportsV6_Nil(t *testing.T) { func TestNodeSupportsV4_ExplicitV4(t *testing.T) { n := &nodeRecord{ServerIPv4: "10.0.0.1"} if !nodeSupportsV4(n) { - t.Fatal("explicit server_ip_v4 must support v4") + t.Fatal("explicit server_ip_v4 needs support v4") } } func TestNodeSupportsV6_ExplicitV6(t *testing.T) { n := &nodeRecord{ServerIPv6: "2001:db8::1"} if !nodeSupportsV6(n) { - t.Fatal("explicit server_ip_v6 must support v6") + t.Fatal("explicit server_ip_v6 needs support v6") } } @@ -68,7 +84,7 @@ func TestNodeSupportsV4_LegacyV4Only(t *testing.T) { t.Fatal("legacy v4 ip in server_ip must support v4") } if nodeSupportsV6(n) { - t.Fatal("legacy v4 ip in server_ip must not support v6") + t.Fatal("legacy v4 ip in server_ip should not support v6") } } @@ -78,7 +94,7 @@ func TestNodeSupportsV6_LegacyV6Only(t *testing.T) { t.Fatal("legacy v6 ip in server_ip must support v6") } if nodeSupportsV4(n) { - t.Fatal("legacy v6 ip in server_ip must not support v4") + t.Fatal("legacy v6 ip in server_ip should not support v4") } } @@ -177,15 +193,15 @@ func v6OnlyNode(name, v6 string) *nodeRecord { } func TestSelectTunnelDialHost_NilNodes(t *testing.T) { - _, err := selectTunnelDialHost(nil, nil, "") + _, err := selectTunnelDialHost(nil, nil, "", "") if err == nil { t.Fatal("expected error for nil nodes") } - _, err = selectTunnelDialHost(dualStackNode("a", "1.1.1.1", "::1"), nil, "") + _, err = selectTunnelDialHost(dualStackNode("a", "1.1.1.1", "::1"), nil, "", "") if err == nil { t.Fatal("expected error for nil toNode") } - _, err = selectTunnelDialHost(nil, dualStackNode("b", "1.1.1.1", "::1"), "") + _, err = selectTunnelDialHost(nil, dualStackNode("b", "1.1.1.1", "::1"), "", "") if err == nil { t.Fatal("expected error for nil fromNode") } @@ -194,8 +210,7 @@ func TestSelectTunnelDialHost_NilNodes(t *testing.T) { func TestSelectTunnelDialHost_DualStack_DefaultPreference(t *testing.T) { from := dualStackNode("from", "10.0.0.1", "2001:db8::1") to := dualStackNode("to", "10.0.0.2", "2001:db8::2") - - host, err := selectTunnelDialHost(from, to, "") + host, err := selectTunnelDialHost(from, to, "", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -208,8 +223,7 @@ func TestSelectTunnelDialHost_DualStack_DefaultPreference(t *testing.T) { func TestSelectTunnelDialHost_DualStack_PreferV4(t *testing.T) { from := dualStackNode("from", "10.0.0.1", "2001:db8::1") to := dualStackNode("to", "10.0.0.2", "2001:db8::2") - - host, err := selectTunnelDialHost(from, to, "v4") + host, err := selectTunnelDialHost(from, to, "v4", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -221,8 +235,7 @@ func TestSelectTunnelDialHost_DualStack_PreferV4(t *testing.T) { func TestSelectTunnelDialHost_DualStack_PreferV6(t *testing.T) { from := dualStackNode("from", "10.0.0.1", "2001:db8::1") to := dualStackNode("to", "10.0.0.2", "2001:db8::2") - - host, err := selectTunnelDialHost(from, to, "v6") + host, err := selectTunnelDialHost(from, to, "v6", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -234,9 +247,8 @@ func TestSelectTunnelDialHost_DualStack_PreferV6(t *testing.T) { func TestSelectTunnelDialHost_V4Only_PreferV6Fallback(t *testing.T) { from := v4OnlyNode("from", "10.0.0.1") to := v4OnlyNode("to", "10.0.0.2") - // User prefers v6, but both nodes are v4-only — should fallback to v4 - host, err := selectTunnelDialHost(from, to, "v6") + host, err := selectTunnelDialHost(from, to, "v6", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -248,9 +260,8 @@ func TestSelectTunnelDialHost_V4Only_PreferV6Fallback(t *testing.T) { func TestSelectTunnelDialHost_V6Only_PreferV4Fallback(t *testing.T) { from := v6OnlyNode("from", "2001:db8::1") to := v6OnlyNode("to", "2001:db8::2") - // User prefers v4, but both nodes are v6-only — should fallback to v6 - host, err := selectTunnelDialHost(from, to, "v4") + host, err := selectTunnelDialHost(from, to, "v4", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -262,8 +273,7 @@ func TestSelectTunnelDialHost_V6Only_PreferV4Fallback(t *testing.T) { func TestSelectTunnelDialHost_Incompatible(t *testing.T) { from := v4OnlyNode("from", "10.0.0.1") to := v6OnlyNode("to", "2001:db8::2") - - _, err := selectTunnelDialHost(from, to, "") + _, err := selectTunnelDialHost(from, to, "", "") if err == nil { t.Fatal("expected error for incompatible nodes (v4-only -> v6-only)") } @@ -272,8 +282,7 @@ func TestSelectTunnelDialHost_Incompatible(t *testing.T) { func TestSelectTunnelDialHost_Incompatible_Reverse(t *testing.T) { from := v6OnlyNode("from", "2001:db8::1") to := v4OnlyNode("to", "10.0.0.2") - - _, err := selectTunnelDialHost(from, to, "") + _, err := selectTunnelDialHost(from, to, "", "") if err == nil { t.Fatal("expected error for incompatible nodes (v6-only -> v4-only)") } @@ -282,9 +291,8 @@ func TestSelectTunnelDialHost_Incompatible_Reverse(t *testing.T) { func TestSelectTunnelDialHost_WhitespacePreference(t *testing.T) { from := dualStackNode("from", "10.0.0.1", "2001:db8::1") to := dualStackNode("to", "10.0.0.2", "2001:db8::2") - // Whitespace should be trimmed, treated as "v6" - host, err := selectTunnelDialHost(from, to, " v6 ") + host, err := selectTunnelDialHost(from, to, " v6 ", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -296,9 +304,8 @@ func TestSelectTunnelDialHost_WhitespacePreference(t *testing.T) { func TestSelectTunnelDialHost_MixedStack_FromDualToV4(t *testing.T) { from := dualStackNode("from", "10.0.0.1", "2001:db8::1") to := v4OnlyNode("to", "10.0.0.2") - // v6 preferred, but target only has v4 — should succeed with v4 - host, err := selectTunnelDialHost(from, to, "v6") + host, err := selectTunnelDialHost(from, to, "v6", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -310,9 +317,8 @@ func TestSelectTunnelDialHost_MixedStack_FromDualToV4(t *testing.T) { func TestSelectTunnelDialHost_MixedStack_FromDualToV6(t *testing.T) { from := dualStackNode("from", "10.0.0.1", "2001:db8::1") to := v6OnlyNode("to", "2001:db8::2") - // v4 preferred, but target only has v6 — should succeed with v6 - host, err := selectTunnelDialHost(from, to, "v4") + host, err := selectTunnelDialHost(from, to, "v4", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -324,9 +330,8 @@ func TestSelectTunnelDialHost_MixedStack_FromDualToV6(t *testing.T) { func TestSelectTunnelDialHost_MixedStack_FromV4ToDual(t *testing.T) { from := v4OnlyNode("from", "10.0.0.1") to := dualStackNode("to", "10.0.0.2", "2001:db8::2") - // v6 preferred, but from only has v4 — should use v4 (from can only reach v4 of target) - host, err := selectTunnelDialHost(from, to, "v6") + host, err := selectTunnelDialHost(from, to, "v6", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -338,9 +343,8 @@ func TestSelectTunnelDialHost_MixedStack_FromV4ToDual(t *testing.T) { func TestSelectTunnelDialHost_MixedStack_FromV6ToDual(t *testing.T) { from := v6OnlyNode("from", "2001:db8::1") to := dualStackNode("to", "10.0.0.2", "2001:db8::2") - // v4 preferred, but from only has v6 — should use v6 - host, err := selectTunnelDialHost(from, to, "v4") + host, err := selectTunnelDialHost(from, to, "v4", "") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -367,7 +371,6 @@ func TestNodeDisplayName_Named(t *testing.T) { t.Fatalf("expected 'hk-node', got %q", got) } } - func TestNodeDisplayName_Unnamed(t *testing.T) { n := &nodeRecord{ID: 42} got := nodeDisplayName(n) diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index aa44704..d314e87 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -271,6 +271,7 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) { nullableText(asString(req["remoteUrl"])), nullableText(asString(req["remoteToken"])), nullableText(asString(req["remoteConfig"])), + nullableText(asString(req["extraIPs"])), ); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -322,6 +323,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) { nullableText(asString(req["serverIpV6"])), defaultString(asString(req["port"]), "1000-65535"), nullableText(asString(req["interfaceName"])), + nullableText(asString(req["extraIPs"])), newHTTP, newTLS, newSocks, @@ -1176,7 +1178,8 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) { if userName == "" { userName = "user" } - forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, nullableInt(speedID)) + inIp := strings.TrimSpace(asString(req["inIp"])) + forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID)) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -1280,6 +1283,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { port = h.pickTunnelPort(tunnelID) } } + inIp := asString(req["inIp"]) fwdEntryNodes, _ := h.tunnelEntryNodeIDs(tunnelID) for _, nodeID := range fwdEntryNodes { node, nodeErr := h.getNodeRecord(nodeID) @@ -1296,7 +1300,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := h.replaceForwardPorts(id, tunnelID, port); err != nil { + if err := h.replaceForwardPorts(id, tunnelID, port, inIp); err != nil { h.rollbackForwardMutation(forward, oldPorts) response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -1628,7 +1632,7 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques fail++ continue } - if err := h.replaceForwardPorts(id, req.TargetTunnelID, p); err != nil { + if err := h.replaceForwardPorts(id, req.TargetTunnelID, p, ""); err != nil { h.rollbackForwardMutation(forward, oldPorts) fail++ continue @@ -1954,6 +1958,7 @@ type tunnelRuntimeNode struct { Inx int ChainType int Port int + ConnectIP string } type tunnelCreateState struct { @@ -2027,6 +2032,7 @@ func (h *Handler) prepareTunnelCreateState(tx *gorm.DB, req map[string]interface Strategy: defaultString(asString(item["strategy"]), "round"), ChainType: 3, Port: port, + ConnectIP: asString(item["connectIp"]), }) } if len(state.OutNodes) == 0 { @@ -2062,6 +2068,7 @@ func (h *Handler) prepareTunnelCreateState(tx *gorm.DB, req map[string]interface Inx: hopIdx + 1, ChainType: 2, Port: port, + ConnectIP: asString(item["connectIp"]), }) } if len(hop) > 0 { @@ -2338,7 +2345,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s h.releaseFederationRuntimeRefs(releaseRefs) return nil, nil, errors.New("节点不存在") } - host, hostErr := selectTunnelDialHost(node, targetNode, state.IPPreference) + host, hostErr := selectTunnelDialHost(node, targetNode, state.IPPreference, target.ConnectIP) if hostErr != nil { h.releaseFederationRuntimeRefs(releaseRefs) return nil, nil, hostErr @@ -2578,7 +2585,7 @@ func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRu if targetNode == nil { return nil, errors.New("节点不存在") } - host, err := selectTunnelDialHost(fromNode, targetNode, ipPreference) + host, err := selectTunnelDialHost(fromNode, targetNode, ipPreference, target.ConnectIP) if err != nil { return nil, err } @@ -2651,10 +2658,13 @@ func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode, return []map[string]interface{}{service} } -func selectTunnelDialHost(fromNode, toNode *nodeRecord, ipPreference string) (string, error) { +func selectTunnelDialHost(fromNode, toNode *nodeRecord, ipPreference string, connectIp string) (string, error) { if fromNode == nil || toNode == nil { return "", errors.New("节点不存在") } + if strings.TrimSpace(connectIp) != "" { + return strings.TrimSpace(connectIp), nil + } fromV4 := nodeSupportsV4(fromNode) fromV6 := nodeSupportsV6(fromNode) toV4 := nodeSupportsV4(toNode) @@ -2789,6 +2799,7 @@ func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[str defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls"), + "", ); err != nil { return err } @@ -2806,6 +2817,7 @@ func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[str return pickErr } } + connectIp := asString(n["connectIp"]) if err := h.repo.CreateChainTunnelTx( tx, tunnelID, @@ -2815,6 +2827,7 @@ func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[str defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls"), + connectIp, ); err != nil { return err } @@ -2834,6 +2847,7 @@ func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[str return pickErr } } + connectIp := asString(n["connectIp"]) if err := h.repo.CreateChainTunnelTx( tx, tunnelID, @@ -2843,6 +2857,7 @@ func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[str defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls"), + connectIp, ); err != nil { return err } @@ -2999,7 +3014,7 @@ func parsePorts(portRange string) ([]int, error) { return ports, nil } -func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int) error { +func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int, inIp string) error { entryNodes, err := h.tunnelEntryNodeIDs(tunnelID) if err != nil { return err @@ -3007,12 +3022,14 @@ func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int) error entries := make([]struct { NodeID int64 Port int + InIP string }, len(entryNodes)) for i, nid := range entryNodes { entries[i] = struct { NodeID int64 Port int - }{NodeID: nid, Port: port} + InIP string + }{NodeID: nid, Port: port, InIP: inIp} } return h.repo.ReplaceForwardPorts(forwardID, entries) } @@ -3021,12 +3038,14 @@ func (h *Handler) replaceForwardPortsWithRecords(forwardID int64, ports []forwar entries := make([]struct { NodeID int64 Port int + InIP string }, len(ports)) for i, fp := range ports { entries[i] = struct { NodeID int64 Port int - }{NodeID: fp.NodeID, Port: fp.Port} + InIP string + }{NodeID: fp.NodeID, Port: fp.Port, InIP: fp.InIP} } return h.repo.ReplaceForwardPorts(forwardID, entries) } diff --git a/go-backend/internal/store/model/model.go b/go-backend/internal/store/model/model.go index fcdc8ca..e485683 100644 --- a/go-backend/internal/store/model/model.go +++ b/go-backend/internal/store/model/model.go @@ -48,10 +48,11 @@ type Forward struct { func (Forward) TableName() string { return "forward" } type ForwardPort struct { - ID int64 `gorm:"primaryKey;autoIncrement"` - ForwardID int64 `gorm:"column:forward_id;not null"` - NodeID int64 `gorm:"column:node_id;not null"` - Port int `gorm:"not null"` + ID int64 `gorm:"primaryKey;autoIncrement"` + ForwardID int64 `gorm:"column:forward_id;not null"` + NodeID int64 `gorm:"column:node_id;not null"` + Port int `gorm:"not null"` + InIP sql.NullString `gorm:"column:in_ip;type:text"` } func (ForwardPort) TableName() string { return "forward_port" } @@ -63,6 +64,7 @@ type Node struct { ServerIP string `gorm:"column:server_ip;type:varchar(100);not null"` ServerIPV4 sql.NullString `gorm:"column:server_ip_v4;type:varchar(100)"` ServerIPV6 sql.NullString `gorm:"column:server_ip_v6;type:varchar(100)"` + ExtraIPs sql.NullString `gorm:"column:extra_ips;type:text"` Port string `gorm:"type:text;not null"` InterfaceName sql.NullString `gorm:"column:interface_name;type:varchar(200)"` Version sql.NullString `gorm:"type:varchar(100)"` @@ -133,6 +135,7 @@ type ChainTunnel struct { Strategy sql.NullString `gorm:"type:varchar(10)"` Inx sql.NullInt64 `gorm:"column:inx"` Protocol sql.NullString `gorm:"type:varchar(10)"` + ConnectIP sql.NullString `gorm:"column:connect_ip;type:varchar(45)"` } func (ChainTunnel) TableName() string { return "chain_tunnel" } @@ -337,6 +340,7 @@ type NodeBackup struct { ServerIP string `json:"serverIp"` ServerIPv4 string `json:"serverIpV4,omitempty"` ServerIPv6 string `json:"serverIpV6,omitempty"` + ExtraIPs string `json:"extraIPs,omitempty"` Port string `json:"port"` InterfaceName string `json:"interfaceName,omitempty"` Version string `json:"version,omitempty"` @@ -510,6 +514,7 @@ type TunnelRecord struct { type ForwardPortRecord struct { NodeID int64 Port int + InIP string } // NodeRecord is a node view used by control plane. @@ -519,6 +524,7 @@ type NodeRecord struct { ServerIP string ServerIPv4 string ServerIPv6 string + ExtraIPs string Status int PortRange string TCPListenAddr string @@ -538,6 +544,7 @@ type ChainNodeRecord struct { NodeName string Protocol string Strategy string + ConnectIP string } type UserTunnelLimiterInfo struct { diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go index 6f42cc1..b2c4b80 100644 --- a/go-backend/internal/store/repo/repository.go +++ b/go-backend/internal/store/repo/repository.go @@ -637,6 +637,7 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) { "ip": n.ServerIP, "serverIp": n.ServerIP, "serverIpV4": nullableString(n.ServerIPV4), "serverIpV6": nullableString(n.ServerIPV6), + "extraIPs": nullableString(n.ExtraIPs), "port": n.Port, "tcpListenAddr": n.TCPListenAddr, "udpListenAddr": n.UDPListenAddr, @@ -2731,10 +2732,11 @@ func resolveForwardIngress(db *gorm.DB, forwardID int64, tunnelID int64) (string type fpRow struct { Port sql.NullInt64 ServerIP sql.NullString + InIP sql.NullString } var fpRows []fpRow err := db.Model(&model.ForwardPort{}). - Select("forward_port.port, node.server_ip"). + Select("forward_port.port, node.server_ip, forward_port.in_ip"). Joins("LEFT JOIN node ON node.id = forward_port.node_id"). Where("forward_port.forward_id = ?", forwardID). Order("forward_port.id ASC"). @@ -2744,10 +2746,22 @@ func resolveForwardIngress(db *gorm.DB, forwardID int64, tunnelID int64) (string } ports := make([]int64, 0) - nodePairs := make([]string, 0) + entries := make([]string, 0) seenPorts := make(map[int64]struct{}) seenPairs := make(map[string]struct{}) + var tunnelFirstIP string + if tunnelInIP.Valid && strings.TrimSpace(tunnelInIP.String) != "" { + tunnelIPs := strings.Split(tunnelInIP.String, ",") + for _, ip := range tunnelIPs { + ip = strings.TrimSpace(ip) + if ip != "" { + tunnelFirstIP = ip + break + } + } + } + for _, row := range fpRows { if !row.Port.Valid { continue @@ -2756,11 +2770,21 @@ func resolveForwardIngress(db *gorm.DB, forwardID int64, tunnelID int64) (string seenPorts[row.Port.Int64] = struct{}{} ports = append(ports, row.Port.Int64) } - if row.ServerIP.Valid && strings.TrimSpace(row.ServerIP.String) != "" { - pair := fmt.Sprintf("%s:%d", strings.TrimSpace(row.ServerIP.String), row.Port.Int64) + + var ip string + if row.InIP.Valid && strings.TrimSpace(row.InIP.String) != "" { + ip = strings.TrimSpace(row.InIP.String) + } else if tunnelFirstIP != "" { + ip = tunnelFirstIP + } else if row.ServerIP.Valid && strings.TrimSpace(row.ServerIP.String) != "" { + ip = strings.TrimSpace(row.ServerIP.String) + } + + if ip != "" { + pair := fmt.Sprintf("%s:%d", ip, row.Port.Int64) if _, ok := seenPairs[pair]; !ok { seenPairs[pair] = struct{}{} - nodePairs = append(nodePairs, pair) + entries = append(entries, pair) } } } @@ -2771,27 +2795,6 @@ func resolveForwardIngress(db *gorm.DB, forwardID int64, tunnelID int64) (string inPort := sql.NullInt64{Int64: ports[0], Valid: true} - entries := make([]string, 0) - if tunnelInIP.Valid && strings.TrimSpace(tunnelInIP.String) != "" { - tunnelIPs := strings.Split(tunnelInIP.String, ",") - seen := make(map[string]struct{}) - for _, ip := range tunnelIPs { - ip = strings.TrimSpace(ip) - if ip == "" { - continue - } - if _, ok := seen[ip]; ok { - continue - } - seen[ip] = struct{}{} - for _, port := range ports { - entries = append(entries, fmt.Sprintf("%s:%d", ip, port)) - } - } - } else { - entries = append(entries, nodePairs...) - } - return strings.Join(entries, ","), inPort, nil } diff --git a/go-backend/internal/store/repo/repository_control.go b/go-backend/internal/store/repo/repository_control.go index 680fae8..2799486 100644 --- a/go-backend/internal/store/repo/repository_control.go +++ b/go-backend/internal/store/repo/repository_control.go @@ -102,7 +102,11 @@ func (r *Repository) ListForwardPorts(forwardID int64) ([]model.ForwardPortRecor } rows := make([]model.ForwardPortRecord, 0, len(ports)) for _, p := range ports { - rows = append(rows, model.ForwardPortRecord{NodeID: p.NodeID, Port: p.Port}) + inIP := "" + if p.InIP.Valid { + inIP = p.InIP.String + } + rows = append(rows, model.ForwardPortRecord{NodeID: p.NodeID, Port: p.Port, InIP: inIP}) } return rows, nil } @@ -177,6 +181,9 @@ func nodeRecordFromModel(n *model.Node) *model.NodeRecord { if n.ServerIPV6.Valid { rec.ServerIPv6 = strings.TrimSpace(n.ServerIPV6.String) } + if n.ExtraIPs.Valid { + rec.ExtraIPs = strings.TrimSpace(n.ExtraIPs.String) + } if n.InterfaceName.Valid { rec.InterfaceName = strings.TrimSpace(n.InterfaceName.String) } @@ -289,10 +296,11 @@ func (r *Repository) ListChainNodesForTunnel(tunnelID int64) ([]model.ChainNodeR Name sql.NullString Protocol sql.NullString Strategy sql.NullString + ConnectIP sql.NullString } var rows []row err := r.db.Model(&model.ChainTunnel{}). - Select("chain_tunnel.chain_type, chain_tunnel.inx, chain_tunnel.node_id, chain_tunnel.port, node.name, chain_tunnel.protocol, chain_tunnel.strategy"). + Select("chain_tunnel.chain_type, chain_tunnel.inx, chain_tunnel.node_id, chain_tunnel.port, node.name, chain_tunnel.protocol, chain_tunnel.strategy, chain_tunnel.connect_ip"). Joins("LEFT JOIN node ON node.id = chain_tunnel.node_id"). Where("chain_tunnel.tunnel_id = ?", tunnelID). Order("chain_tunnel.chain_type ASC, chain_tunnel.inx ASC, chain_tunnel.id ASC"). @@ -337,6 +345,9 @@ func (r *Repository) ListChainNodesForTunnel(tunnelID int64) ([]model.ChainNodeR } else { item.Strategy = row.Strategy.String } + if row.ConnectIP.Valid { + item.ConnectIP = row.ConnectIP.String + } result = append(result, item) } return result, nil diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go index c7fb53b..76cc6bc 100644 --- a/go-backend/internal/store/repo/repository_mutations.go +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -196,7 +196,7 @@ func (r *Repository) GetUserDefaultsForTunnel(userID int64) (flow int64, num int return user.Flow, user.Num, user.ExpTime, user.FlowResetTime, nil } -func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig interface{}) error { +func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}) error { if r == nil || r.db == nil { return errors.New("repository not initialized") } @@ -206,6 +206,7 @@ func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serve ServerIP: serverIP, ServerIPV4: nullStringFromInterface(serverIPV4), ServerIPV6: nullStringFromInterface(serverIPV6), + ExtraIPs: nullStringFromInterface(extraIPs), Port: stringFromInterface(port), InterfaceName: nullStringFromInterface(interfaceName), Version: nullStringFromInterface(version), @@ -238,7 +239,7 @@ func (r *Repository) GetNodeStatusFields(nodeID int64) (status, httpFlag, tlsFla return node.Status, node.HTTP, node.TLS, node.Socks, nil } -func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error { +func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error { if r == nil || r.db == nil { return errors.New("repository not initialized") } @@ -249,6 +250,7 @@ func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, ser "server_ip": serverIP, "server_ip_v4": nullStringFromInterface(serverIPV4), "server_ip_v6": nullStringFromInterface(serverIPV6), + "extra_ips": nullStringFromInterface(extraIPs), "port": stringFromInterface(port), "interface_name": nullStringFromInterface(interfaceName), "http": httpFlag, @@ -395,7 +397,7 @@ func (r *Repository) DeleteChainTunnelsByTunnelTx(tx *gorm.DB, tunnelID int64) e return tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error } -func (r *Repository) CreateChainTunnelTx(tx *gorm.DB, tunnelID int64, chainType string, nodeID int64, port sql.NullInt64, strategy string, inx int, protocol string) error { +func (r *Repository) CreateChainTunnelTx(tx *gorm.DB, tunnelID int64, chainType string, nodeID int64, port sql.NullInt64, strategy string, inx int, protocol string, connectIp string) error { if tx == nil { return errors.New("database unavailable") } @@ -407,6 +409,7 @@ func (r *Repository) CreateChainTunnelTx(tx *gorm.DB, tunnelID int64, chainType Strategy: nullStringFromInterface(strategy), Inx: nullInt64FromInterface(inx), Protocol: nullStringFromInterface(protocol), + ConnectIP: sql.NullString{String: connectIp, Valid: connectIp != ""}, } return tx.Create(&ct).Error } @@ -692,6 +695,7 @@ func (r *Repository) DeleteForwardCascade(forwardID int64) error { func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct { NodeID int64 Port int + InIP string }) error { if r == nil || r.db == nil { return errors.New("repository not initialized") @@ -705,7 +709,12 @@ func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct { } rows := make([]model.ForwardPort, 0, len(entries)) for _, e := range entries { - rows = append(rows, model.ForwardPort{ForwardID: forwardID, NodeID: e.NodeID, Port: e.Port}) + rows = append(rows, model.ForwardPort{ + ForwardID: forwardID, + NodeID: e.NodeID, + Port: e.Port, + InIP: sql.NullString{String: e.InIP, Valid: e.InIP != ""}, + }) } return tx.Create(&rows).Error }) @@ -1168,7 +1177,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool, return ut.ID, true, nil } -func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, speedID interface{}) (int64, error) { +func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}) (int64, error) { if r == nil || r.db == nil { return 0, errors.New("repository not initialized") } @@ -1198,6 +1207,7 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel ForwardID: forwardID, NodeID: nodeID, Port: port, + InIP: sql.NullString{String: inIp, Valid: inIp != ""}, } if err := tx.Create(&fp).Error; err != nil { return err diff --git a/vite-frontend/src/pages/forward.tsx b/vite-frontend/src/pages/forward.tsx index b9c94fd..c2816ef 100644 --- a/vite-frontend/src/pages/forward.tsx +++ b/vite-frontend/src/pages/forward.tsx @@ -127,6 +127,7 @@ interface ForwardForm { name: string; tunnelId: number | null; inPort: number | null; + inIp: string; remoteAddr: string; interfaceName?: string; strategy: string; @@ -546,6 +547,7 @@ export default function ForwardPage() { name: "", tunnelId: null, inPort: null, + inIp: "", remoteAddr: "", interfaceName: "", strategy: "fifo", @@ -1164,6 +1166,7 @@ export default function ForwardPage() { name: "", tunnelId: null, inPort: null, + inIp: "", remoteAddr: "", interfaceName: "", strategy: "fifo", @@ -1182,6 +1185,7 @@ export default function ForwardPage() { name: forward.name, tunnelId: forward.tunnelId, inPort: forward.inPort, + inIp: forward.inIp || "", remoteAddr: forward.remoteAddr.split(",").join("\n"), interfaceName: forward.interfaceName || "", strategy: forward.strategy || "fifo", @@ -1263,6 +1267,7 @@ export default function ForwardPage() { name: form.name, tunnelId: form.tunnelId, inPort: form.inPort, + inIp: form.inIp || null, remoteAddr: processedRemoteAddr, strategy: addressCount > 1 ? form.strategy : "fifo", speedId: normalizeSpeedId(form.speedId), @@ -1270,11 +1275,11 @@ export default function ForwardPage() { res = await updateForward(updateData); } else { - // 创建时不需要id和userId(后端会自动设置) const createData = { name: form.name, tunnelId: form.tunnelId, inPort: form.inPort, + inIp: form.inIp || null, remoteAddr: processedRemoteAddr, strategy: addressCount > 1 ? form.strategy : "fifo", speedId: normalizeSpeedId(form.speedId), @@ -4023,6 +4028,17 @@ export default function ForwardPage() { }} /> + + setForm((prev) => ({ ...prev, inIp: e.target.value })) + } + /> +