From e38335973d8338c4431e9407bdd99c46be9df312 Mon Sep 17 00:00:00 2001 From: sagitchu Date: Tue, 3 Mar 2026 10:58:15 +0800 Subject: [PATCH] fix: apply custom IP binding to forward and tunnel chain services Entire-Checkpoint: ceff329d4cf4 --- .../internal/http/handler/control_plane.go | 9 ++++-- .../http/handler/control_plane_test.go | 32 +++++++++++++++++++ .../internal/http/handler/dual_stack_test.go | 26 +++++++++++++++ go-backend/internal/http/handler/mutations.go | 2 +- 4 files changed, 65 insertions(+), 4 deletions(-) diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 5dc4b33..0b4b2de 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -281,7 +281,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all if err != nil { return err } - services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiterID, tunnelTLSProtocol) + services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, tunnelTLSProtocol) _, err = h.sendNodeCommand(node.ID, method, services, true, false) if err != nil && allowFallbackAdd && method == "UpdateService" { _, err = h.sendNodeCommand(node.ID, "AddService", services, true, false) @@ -1303,7 +1303,7 @@ func isAlreadyExistsMessage(message string) bool { return strings.Contains(msg, "already exists") || strings.Contains(msg, "已存在") } -func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiterID *int64, tunnelTLSProtocol bool) []map[string]interface{} { +func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64, tunnelTLSProtocol bool) []map[string]interface{} { protocols := []string{"tcp", "udp"} services := make([]map[string]interface{}, 0, 2) targets := splitRemoteTargets(forward.RemoteAddr) @@ -1317,9 +1317,12 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel if protocol == "udp" { listenerAddr = node.UDPListenAddr } + if bindIP != "" { + listenerAddr = bindIP + } service := map[string]interface{}{ "name": fmt.Sprintf("%s_%s", baseName, protocol), - "addr": fmt.Sprintf("%s:%d", listenerAddr, port), + "addr": processServerAddress(fmt.Sprintf("%s:%d", listenerAddr, port)), "handler": map[string]interface{}{ "type": protocol, }, diff --git a/go-backend/internal/http/handler/control_plane_test.go b/go-backend/internal/http/handler/control_plane_test.go index 963d7a9..a09d8fa 100644 --- a/go-backend/internal/http/handler/control_plane_test.go +++ b/go-backend/internal/http/handler/control_plane_test.go @@ -65,3 +65,35 @@ func TestIsAlreadyExistsMessage(t *testing.T) { t.Fatalf("address already in use must not be treated as already exists") } } + +func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) { + forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7} + node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"} + services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil, false) + if len(services) != 2 { + t.Fatalf("expected 2 services, got %d", len(services)) + } + for _, svc := range services { + addr, _ := svc["addr"].(string) + if addr != "10.9.8.7:22000" { + t.Fatalf("expected bind IP address 10.9.8.7:22000, got %q", addr) + } + } +} + +func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) { + forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7} + node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"} + services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil, false) + if len(services) != 2 { + t.Fatalf("expected 2 services, got %d", len(services)) + } + tcpAddr, _ := services[0]["addr"].(string) + udpAddr, _ := services[1]["addr"].(string) + if tcpAddr != "0.0.0.0:22001" { + t.Fatalf("expected tcp addr 0.0.0.0:22001, got %q", tcpAddr) + } + if udpAddr != "[::]:22001" { + t.Fatalf("expected udp addr [::]:22001, got %q", udpAddr) + } +} diff --git a/go-backend/internal/http/handler/dual_stack_test.go b/go-backend/internal/http/handler/dual_stack_test.go index 3c423a4..476f1ce 100644 --- a/go-backend/internal/http/handler/dual_stack_test.go +++ b/go-backend/internal/http/handler/dual_stack_test.go @@ -30,6 +30,32 @@ func TestSelectTunnelDialHost_ConnectIpPriority(t *testing.T) { } } +func TestBuildTunnelChainServiceConfig_UsesConnectIPForListen(t *testing.T) { + node := &nodeRecord{TCPListenAddr: "[::]"} + chain := tunnelRuntimeNode{Protocol: "tls", Port: 21000, ConnectIP: "2001:db8::88"} + services := buildTunnelChainServiceConfig(99, chain, node) + if len(services) != 1 { + t.Fatalf("expected 1 service, got %d", len(services)) + } + addr, _ := services[0]["addr"].(string) + if addr != "[2001:db8::88]:21000" { + t.Fatalf("expected connectIp listen [2001:db8::88]:21000, got %q", addr) + } +} + +func TestBuildTunnelChainServiceConfig_DefaultListenAddrWhenConnectIPEmpty(t *testing.T) { + node := &nodeRecord{TCPListenAddr: "[::]"} + chain := tunnelRuntimeNode{Protocol: "tls", Port: 21001} + services := buildTunnelChainServiceConfig(99, chain, node) + if len(services) != 1 { + t.Fatalf("expected 1 service, got %d", len(services)) + } + addr, _ := services[0]["addr"].(string) + if addr != "[::]:21001" { + t.Fatalf("expected default listen [::]:21001, got %q", addr) + } +} + func TestNodeSupportsV6_Nil(t *testing.T) { if nodeSupportsV6(nil) { t.Fatal("nil node must not support v6") diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index d314e87..6f3b431 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -2643,7 +2643,7 @@ func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode, } service := map[string]interface{}{ "name": fmt.Sprintf("%d_tls", tunnelID), - "addr": fmt.Sprintf("%s:%d", node.TCPListenAddr, chainNode.Port), + "addr": processServerAddress(fmt.Sprintf("%s:%d", defaultString(strings.TrimSpace(chainNode.ConnectIP), node.TCPListenAddr), chainNode.Port)), "handler": handlerCfg, "listener": map[string]interface{}{ "type": protocol,