fix: use configured connectIp for tunnel chain diagnosis

- Pass connectIp through resolveChainProbeTarget in diagnosis stream start items
- Pass connectIp in appendChainHopDiagnosis for full chain probes
- Reconstruct tunnel state with connectIp field preserved
- Fix forward service config when bindIP already contains port
- Add contract tests for connectIp diagnosis scenarios
- Add unit test for bindIP with port in buildForwardServiceConfigs
- Update AGENTS.md with plan document rules

Entire-Checkpoint: 35a2e61c2431
This commit is contained in:
sagitchu
2026-03-03 14:17:36 +08:00
parent e38335973d
commit 1c10347357
7 changed files with 295 additions and 5 deletions
@@ -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, workItem.toNode.ConnectIP)
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, toNode.ConnectIP)
if err != nil {
h.appendFailedDiagnosis(results, nodeCache, fromNodeID, strings.Trim(strings.TrimSpace(targetNode.ServerIP), "[]"), toNode.Port, description, metadata, err.Error())
return
@@ -1317,12 +1317,19 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
if protocol == "udp" {
listenerAddr = node.UDPListenAddr
}
var serviceAddr string
if bindIP != "" {
listenerAddr = bindIP
if strings.Contains(bindIP, ":") {
serviceAddr = processServerAddress(bindIP)
} else {
serviceAddr = processServerAddress(fmt.Sprintf("%s:%d", bindIP, port))
}
} else {
serviceAddr = processServerAddress(fmt.Sprintf("%s:%d", listenerAddr, port))
}
service := map[string]interface{}{
"name": fmt.Sprintf("%s_%s", baseName, protocol),
"addr": processServerAddress(fmt.Sprintf("%s:%d", listenerAddr, port)),
"addr": serviceAddr,
"handler": map[string]interface{}{
"type": protocol,
},
@@ -97,3 +97,18 @@ func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.
t.Fatalf("expected udp addr [::]:22001, got %q", udpAddr)
}
}
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(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, 55555, "3.3.3.3:12345", 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 != "3.3.3.3:12345" {
t.Fatalf("expected bind IP with port 3.3.3.3:12345, got %q", addr)
}
}
}
@@ -887,6 +887,7 @@ func (h *Handler) reconstructTunnelState(tunnelID int64) (*tunnelCreateState, er
Strategy: r.Strategy,
ChainType: 3,
Port: r.Port,
ConnectIP: r.ConnectIP,
})
state.NodeIDList = append(state.NodeIDList, r.NodeID)
}
@@ -901,6 +902,7 @@ func (h *Handler) reconstructTunnelState(tunnelID int64) (*tunnelCreateState, er
ChainType: 2,
Inx: int(r.Inx),
Port: r.Port,
ConnectIP: r.ConnectIP,
})
state.NodeIDList = append(state.NodeIDList, r.NodeID)
}
@@ -0,0 +1,79 @@
package handler
import (
"path/filepath"
"testing"
"time"
"go-backend/internal/store/repo"
)
func TestReconstructTunnelState_PreservesConnectIP(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "reconstruct-connect-ip.db")
r, err := repo.Open(dbPath)
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = r.Close() })
h := New(r, "secret")
now := time.Now().UnixMilli()
if err := r.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(1, 'reconstruct-tunnel', 1.0, 2, 'tls', 1, ?, ?, 1, NULL, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
insertNode := func(id int64, name, ip string) {
if err := r.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, id, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node %s: %v", name, err)
}
}
insertNode(101, "entry", "10.90.0.10")
insertNode(102, "middle", "10.90.0.20")
insertNode(103, "exit", "10.90.0.30")
if err := r.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(1, '1', 101, 30001, 'round', 1, 'tls')
`).Error; err != nil {
t.Fatalf("insert entry chain: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol, connect_ip)
VALUES(1, '2', 102, 30002, 'round', 1, 'tls', '10.99.9.22')
`).Error; err != nil {
t.Fatalf("insert middle chain: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol, connect_ip)
VALUES(1, '3', 103, 30003, 'round', 1, 'tls', '10.99.9.33')
`).Error; err != nil {
t.Fatalf("insert exit chain: %v", err)
}
state, err := h.reconstructTunnelState(1)
if err != nil {
t.Fatalf("reconstructTunnelState: %v", err)
}
if len(state.ChainHops) != 1 || len(state.ChainHops[0]) != 1 {
t.Fatalf("unexpected chain hops: %+v", state.ChainHops)
}
if got := state.ChainHops[0][0].ConnectIP; got != "10.99.9.22" {
t.Fatalf("expected middle connectIp 10.99.9.22, got %q", got)
}
if len(state.OutNodes) != 1 {
t.Fatalf("unexpected out nodes: %+v", state.OutNodes)
}
if got := state.OutNodes[0].ConnectIP; got != "10.99.9.33" {
t.Fatalf("expected exit connectIp 10.99.9.33, got %q", got)
}
}