From f887deb44cccceb2da40e31d1efcba33ff05d420 Mon Sep 17 00:00:00 2001 From: mango <1711456624@qq.com> Date: Sun, 9 Aug 2026 14:20:33 +0800 Subject: [PATCH] =?UTF-8?q?=F0=9F=90=9B=20fix(rocketmq):=20=E6=94=AF?= =?UTF-8?q?=E6=8C=81=E7=BB=8F=E9=9A=A7=E9=81=93=E8=AE=BF=E9=97=AE=E9=9A=94?= =?UTF-8?q?=E7=A6=BB=E7=BD=91=E7=BB=9C=20(#886)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #886 --- internal/app/db_proxy.go | 5 + internal/app/db_proxy_test.go | 69 ++++ internal/db/rocketmq_impl.go | 18 +- internal/db/rocketmq_impl_test.go | 84 +++++ internal/db/rocketmq_tunnel.go | 534 ++++++++++++++++++++++++++++ internal/db/rocketmq_tunnel_test.go | 500 ++++++++++++++++++++++++++ 6 files changed, 1205 insertions(+), 5 deletions(-) create mode 100644 internal/db/rocketmq_tunnel.go create mode 100644 internal/db/rocketmq_tunnel_test.go diff --git a/internal/app/db_proxy.go b/internal/app/db_proxy.go index 658e94a0..9f2c989a 100644 --- a/internal/app/db_proxy.go +++ b/internal/app/db_proxy.go @@ -88,6 +88,11 @@ func resolveDialConfigWithProxy(raw connection.ConnectionConfig) (connection.Con // normalized proxy directly instead of using a local TCP forwarder. return config, nil } + if normalizedType == "rocketmq" || normalizedType == "rocket-mq" || normalizedType == "rocket_mq" || normalizedType == "apache-rocketmq" || normalizedType == "apache_rocketmq" || normalizedType == "rmq" { + // RocketMQ discovers broker addresses from NameServer responses. Its + // driver must keep the proxy so it can forward both address layers. + return config, nil + } if normalizedType == "sqlite" || normalizedType == "duckdb" || normalizedType == "custom" { // 文件型/自定义 DSN 类型不走标准 host:port,不在此层改写。 return config, nil diff --git a/internal/app/db_proxy_test.go b/internal/app/db_proxy_test.go index 878b2278..67c98c63 100644 --- a/internal/app/db_proxy_test.go +++ b/internal/app/db_proxy_test.go @@ -186,3 +186,72 @@ func TestResolveDialConfigWithProxy_NacosProxyWithSSHForwardsGatewayOnly(t *test t.Fatalf("proxy should only wrap the SSH gateway, got enabled=%v config=%#v", got.UseProxy, got.Proxy) } } + +func TestResolveDialConfigWithProxy_RocketMQKeepsDynamicTargets(t *testing.T) { + tests := []struct { + name string + raw connection.ConnectionConfig + want connection.ProxyConfig + }{ + { + name: "explicit socks5 proxy", + raw: connection.ConnectionConfig{ + Type: "rocketmq", + Host: "nameserver.internal.test", + Port: 9876, + Hosts: []string{"nameserver-backup.internal.test:9876"}, + UseProxy: true, + Proxy: connection.ProxyConfig{ + Type: "socks5h", + Host: "127.0.0.1", + Port: 1080, + }, + }, + want: connection.ProxyConfig{ + Type: "socks5", + Host: "127.0.0.1", + Port: 1080, + }, + }, + { + name: "HTTP tunnel", + raw: connection.ConnectionConfig{ + Type: "rocketmq", + Host: "nameserver.internal.test", + Port: 9876, + UseHTTPTunnel: true, + HTTPTunnel: connection.HTTPTunnelConfig{ + Host: "tunnel.internal.test", + Port: 8080, + User: "tunnel-user", + Password: "tunnel-password", + }, + }, + want: connection.ProxyConfig{ + Type: "http", + Host: "tunnel.internal.test", + Port: 8080, + User: "tunnel-user", + Password: "tunnel-password", + }, + }, + } + + for _, testCase := range tests { + t.Run(testCase.name, func(t *testing.T) { + got, err := resolveDialConfigWithProxy(testCase.raw) + if err != nil { + t.Fatalf("resolveDialConfigWithProxy: %v", err) + } + if got.Host != testCase.raw.Host || got.Port != testCase.raw.Port || !reflect.DeepEqual(got.Hosts, testCase.raw.Hosts) { + t.Fatalf("RocketMQ targets = %s:%d %v, want %s:%d %v", got.Host, got.Port, got.Hosts, testCase.raw.Host, testCase.raw.Port, testCase.raw.Hosts) + } + if !got.UseProxy || got.Proxy != testCase.want { + t.Fatalf("RocketMQ proxy = %#v (enabled=%v), want %#v", got.Proxy, got.UseProxy, testCase.want) + } + if got.UseHTTPTunnel || got.HTTPTunnel != (connection.HTTPTunnelConfig{}) { + t.Fatalf("HTTP tunnel was not normalized into proxy config: %#v", got.HTTPTunnel) + } + }) + } +} diff --git a/internal/db/rocketmq_impl.go b/internal/db/rocketmq_impl.go index b091e57c..a5685d5a 100644 --- a/internal/db/rocketmq_impl.go +++ b/internal/db/rocketmq_impl.go @@ -125,6 +125,7 @@ var newRocketMQRuntime = func(config connection.ConnectionConfig) (rocketmqRunti type RocketMQDB struct { runtime rocketmqRuntime + tunnel *rocketmqTunnelSet defaultTopic string defaultConsumerGroup string defaultTagExpression string @@ -137,15 +138,16 @@ func (r *RocketMQDB) Connect(config connection.ConnectionConfig) error { _ = r.Close() runConfig := normalizeRocketMQConfig(config) - if runConfig.UseSSH { - return fmt.Errorf("RocketMQ 当前暂不支持 SSH 隧道;请直接连通 NameServer 与 Broker") - } - if runConfig.UseProxy || runConfig.UseHTTPTunnel { - return fmt.Errorf("RocketMQ 当前暂不支持代理或 HTTP 隧道;请直接连通 NameServer 与 Broker") + preparedConfig, tunnel, err := prepareRocketMQTunnel(runConfig) + if err != nil { + return err } + r.tunnel = tunnel + runConfig = preparedConfig runtime, err := newRocketMQRuntime(runConfig) if err != nil { + _ = r.Close() return err } r.runtime = runtime @@ -170,7 +172,13 @@ func (r *RocketMQDB) Close() error { firstErr = err } } + if r.tunnel != nil { + if err := r.tunnel.Close(); err != nil && firstErr == nil { + firstErr = err + } + } r.runtime = nil + r.tunnel = nil r.defaultTopic = "" r.defaultConsumerGroup = "" r.defaultTagExpression = "" diff --git a/internal/db/rocketmq_impl_test.go b/internal/db/rocketmq_impl_test.go index f719aee4..456235aa 100644 --- a/internal/db/rocketmq_impl_test.go +++ b/internal/db/rocketmq_impl_test.go @@ -2,6 +2,7 @@ package db import ( "context" + "net" "reflect" "strings" "testing" @@ -90,6 +91,89 @@ func TestNormalizeRocketMQConfigParsesURIAndParams(t *testing.T) { } } +func TestRocketMQConnectSupportsNetworkTunnels(t *testing.T) { + tests := []struct { + name string + config connection.ConnectionConfig + }{ + { + name: "SSH", + config: connection.ConnectionConfig{ + Type: "rocketmq", + Host: "nameserver.internal.test", + Port: 9876, + UseSSH: true, + SSH: connection.SSHConfig{ + Host: "ssh.internal.test", + Port: 22, + User: "ssh-user", + }, + }, + }, + { + name: "proxy", + config: connection.ConnectionConfig{ + Type: "rocketmq", + Host: "nameserver.internal.test", + Port: 9876, + UseProxy: true, + Proxy: connection.ProxyConfig{ + Type: "socks5", + Host: "proxy.internal.test", + Port: 1080, + }, + }, + }, + { + name: "HTTP tunnel", + config: connection.ConnectionConfig{ + Type: "rocketmq", + Host: "nameserver.internal.test", + Port: 9876, + UseHTTPTunnel: true, + HTTPTunnel: connection.HTTPTunnelConfig{ + Host: "tunnel.internal.test", + Port: 8080, + }, + }, + }, + } + + for _, testCase := range tests { + t.Run(testCase.name, func(t *testing.T) { + var runtimeConfig connection.ConnectionConfig + originalFactory := newRocketMQRuntime + newRocketMQRuntime = func(config connection.ConnectionConfig) (rocketmqRuntime, error) { + runtimeConfig = config + return &fakeRocketMQRuntime{}, nil + } + defer func() { + newRocketMQRuntime = originalFactory + }() + + client := &RocketMQDB{} + if err := client.Connect(testCase.config); err != nil { + t.Fatalf("Connect failed: %v", err) + } + + if runtimeConfig.UseSSH || runtimeConfig.UseProxy || runtimeConfig.UseHTTPTunnel { + t.Fatalf("runtime received unresolved tunnel config: %#v", runtimeConfig) + } + if runtimeConfig.Host != "127.0.0.1" || runtimeConfig.Port <= 0 { + t.Fatalf("runtime NameServer = %s:%d, want local forwarded address", runtimeConfig.Host, runtimeConfig.Port) + } + localAddress := rocketmqFormatHostPort(runtimeConfig.Host, runtimeConfig.Port) + if err := client.Close(); err != nil { + t.Fatalf("Close failed: %v", err) + } + if conn, err := net.DialTimeout("tcp", localAddress, 50*time.Millisecond); err == nil { + _ = conn.Close() + t.Fatalf("RocketMQ tunnel listener still accepts connections after Close: %s", localAddress) + } + }) + } +} + func TestRocketMQQueryExecAndColumns(t *testing.T) { fakeRuntime := &fakeRocketMQRuntime{ listTopicsResult: []rocketmqTopicInfo{ diff --git a/internal/db/rocketmq_tunnel.go b/internal/db/rocketmq_tunnel.go new file mode 100644 index 00000000..f8ee5504 --- /dev/null +++ b/internal/db/rocketmq_tunnel.go @@ -0,0 +1,534 @@ +package db + +import ( + "context" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "strings" + "sync" + "time" + + "GoNavi-Wails/internal/connection" + "GoNavi-Wails/internal/logger" + proxytunnel "GoNavi-Wails/internal/proxy" + "GoNavi-Wails/internal/ssh" +) + +const ( + defaultRocketMQBrokerPort = 10911 + maxRocketMQFrameSize = 64 << 20 + rocketMQBrokerForwarderIdle = 10 * time.Minute +) + +var ( + rocketMQTunnelNow = time.Now + rocketMQDialContextThroughSSH = ssh.DialContextThroughSSH +) + +type rocketmqDialContextFunc func(ctx context.Context, network, address string) (net.Conn, error) + +type rocketmqTunnelSet struct { + dialContext rocketmqDialContextFunc + + mu sync.Mutex + closed bool + forwarders map[*rocketmqForwarder]struct{} + brokerForwarders map[string]rocketmqBrokerForwarder +} + +type rocketmqBrokerForwarder struct { + forwarder *rocketmqForwarder + lastSeen time.Time +} + +type rocketmqForwarder struct { + listener net.Listener + remoteAddr string + dialContext rocketmqDialContextFunc + rewrite func([]byte) ([]byte, error) + ctx context.Context + cancel context.CancelFunc + + mu sync.Mutex + closed bool + active map[net.Conn]struct{} + closeOnce sync.Once + wg sync.WaitGroup +} + +func prepareRocketMQTunnel(config connection.ConnectionConfig) (connection.ConnectionConfig, *rocketmqTunnelSet, error) { + runConfig, err := normalizeRocketMQTunnelConfig(config) + if err != nil { + return connection.ConnectionConfig{}, nil, err + } + if !runConfig.UseSSH && !runConfig.UseProxy { + return runConfig, nil, nil + } + + var dialContext rocketmqDialContextFunc + switch { + case runConfig.UseSSH: + sshConfig := runConfig.SSH + dialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + return rocketMQDialContextThroughSSH(ctx, sshConfig, network, address) + } + case runConfig.UseProxy: + proxyConfig := runConfig.Proxy + dialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + return proxytunnel.DialContext(ctx, proxyConfig, network, address) + } + } + + nameservers, err := rocketmqNameServerAddresses(runConfig) + if err != nil { + return connection.ConnectionConfig{}, nil, err + } + tunnels := newRocketMQTunnelSet(dialContext) + forwarded := make([]string, 0, len(nameservers)) + for _, nameserver := range nameservers { + localAddr, forwardErr := tunnels.forwardNameServer(nameserver) + if forwardErr != nil { + _ = tunnels.Close() + return connection.ConnectionConfig{}, nil, fmt.Errorf("创建 RocketMQ NameServer 隧道失败:%w", forwardErr) + } + forwarded = append(forwarded, localAddr) + } + + host, port, ok := parseHostPortWithDefault(forwarded[0], 0) + if !ok { + _ = tunnels.Close() + return connection.ConnectionConfig{}, nil, fmt.Errorf("解析 RocketMQ 本地 NameServer 地址失败:%s", forwarded[0]) + } + runConfig.Host = host + runConfig.Port = port + runConfig.Hosts = append([]string(nil), forwarded[1:]...) + runConfig.UseSSH = false + runConfig.SSH = connection.SSHConfig{} + runConfig.UseProxy = false + runConfig.Proxy = connection.ProxyConfig{} + runConfig.UseHTTPTunnel = false + runConfig.HTTPTunnel = connection.HTTPTunnelConfig{} + return runConfig, tunnels, nil +} + +func normalizeRocketMQTunnelConfig(config connection.ConnectionConfig) (connection.ConnectionConfig, error) { + if config.UseHTTPTunnel { + if config.UseProxy { + return connection.ConnectionConfig{}, fmt.Errorf("RocketMQ 不能同时启用代理和 HTTP 隧道") + } + host := strings.TrimSpace(config.HTTPTunnel.Host) + if host == "" { + return connection.ConnectionConfig{}, fmt.Errorf("RocketMQ HTTP 隧道主机不能为空") + } + port := config.HTTPTunnel.Port + if port <= 0 { + port = 8080 + } + if port > 65535 { + return connection.ConnectionConfig{}, fmt.Errorf("RocketMQ HTTP 隧道端口无效:%d", config.HTTPTunnel.Port) + } + config.UseProxy = true + config.Proxy = connection.ProxyConfig{ + Type: "http", + Host: host, + Port: port, + User: strings.TrimSpace(config.HTTPTunnel.User), + Password: config.HTTPTunnel.Password, + } + config.UseHTTPTunnel = false + config.HTTPTunnel = connection.HTTPTunnelConfig{} + } + if config.UseSSH && config.SSH.Port <= 0 { + config.SSH.Port = 22 + } + if config.UseSSH && config.UseProxy { + return connection.ConnectionConfig{}, fmt.Errorf("RocketMQ 同时使用 SSH 和代理时,代理只能用于连接 SSH 网关") + } + if config.UseProxy { + proxyConfig, err := proxytunnel.NormalizeConfig(config.Proxy) + if err != nil { + return connection.ConnectionConfig{}, err + } + config.Proxy = proxyConfig + } + return config, nil +} + +func newRocketMQTunnelSet(dialContext rocketmqDialContextFunc) *rocketmqTunnelSet { + return &rocketmqTunnelSet{ + dialContext: dialContext, + forwarders: make(map[*rocketmqForwarder]struct{}), + brokerForwarders: make(map[string]rocketmqBrokerForwarder), + } +} + +func (t *rocketmqTunnelSet) forwardNameServer(remoteAddr string) (string, error) { + forwarder, err := newRocketMQForwarder(remoteAddr, t.dialContext, func(frame []byte) ([]byte, error) { + return rewriteRocketMQRouteFrame(frame, t.forwardBrokerAddress) + }) + if err != nil { + return "", err + } + + t.mu.Lock() + if t.closed { + t.mu.Unlock() + _ = forwarder.Close() + return "", fmt.Errorf("RocketMQ 隧道已关闭") + } + t.forwarders[forwarder] = struct{}{} + t.mu.Unlock() + return forwarder.LocalAddr(), nil +} + +func (t *rocketmqTunnelSet) forwardBrokerAddress(remoteAddr string) (string, error) { + host, port, ok := parseHostPortWithDefault(remoteAddr, defaultRocketMQBrokerPort) + if !ok { + return "", fmt.Errorf("解析 RocketMQ Broker 地址失败:%s", remoteAddr) + } + canonical := rocketmqFormatHostPort(host, port) + now := rocketMQTunnelNow() + + t.mu.Lock() + if t.closed { + t.mu.Unlock() + return "", fmt.Errorf("RocketMQ 隧道已关闭") + } + if existing, exists := t.brokerForwarders[canonical]; exists { + existing.lastSeen = now + t.brokerForwarders[canonical] = existing + expired := t.collectExpiredBrokerForwardersLocked(now) + t.mu.Unlock() + for _, stale := range expired { + _ = stale.Close() + } + return existing.forwarder.LocalAddr(), nil + } + forwarder, err := newRocketMQForwarder(canonical, t.dialContext, nil) + if err != nil { + t.mu.Unlock() + return "", err + } + t.forwarders[forwarder] = struct{}{} + t.brokerForwarders[canonical] = rocketmqBrokerForwarder{forwarder: forwarder, lastSeen: now} + expired := t.collectExpiredBrokerForwardersLocked(now) + t.mu.Unlock() + + for _, stale := range expired { + _ = stale.Close() + } + logger.Infof("已映射 RocketMQ Broker:本地 %s -> 远端 %s", forwarder.LocalAddr(), canonical) + return forwarder.LocalAddr(), nil +} + +func (t *rocketmqTunnelSet) collectExpiredBrokerForwardersLocked(now time.Time) []*rocketmqForwarder { + expired := make([]*rocketmqForwarder, 0) + for address, entry := range t.brokerForwarders { + if now.Sub(entry.lastSeen) < rocketMQBrokerForwarderIdle || entry.forwarder.HasActiveConnections() { + continue + } + delete(t.brokerForwarders, address) + delete(t.forwarders, entry.forwarder) + expired = append(expired, entry.forwarder) + } + return expired +} + +func (t *rocketmqTunnelSet) Close() error { + if t == nil { + return nil + } + t.mu.Lock() + if t.closed { + t.mu.Unlock() + return nil + } + t.closed = true + forwarders := make([]*rocketmqForwarder, 0, len(t.forwarders)) + for forwarder := range t.forwarders { + forwarders = append(forwarders, forwarder) + } + t.forwarders = nil + t.brokerForwarders = nil + t.mu.Unlock() + + var firstErr error + for _, forwarder := range forwarders { + if err := forwarder.Close(); err != nil && firstErr == nil { + firstErr = err + } + } + return firstErr +} + +func newRocketMQForwarder(remoteAddr string, dialContext rocketmqDialContextFunc, rewrite func([]byte) ([]byte, error)) (*rocketmqForwarder, error) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + return nil, err + } + ctx, cancel := context.WithCancel(context.Background()) + forwarder := &rocketmqForwarder{ + listener: listener, + remoteAddr: remoteAddr, + dialContext: dialContext, + rewrite: rewrite, + ctx: ctx, + cancel: cancel, + active: make(map[net.Conn]struct{}), + } + forwarder.wg.Add(1) + go forwarder.serve() + return forwarder, nil +} + +func (f *rocketmqForwarder) LocalAddr() string { + return f.listener.Addr().String() +} + +func (f *rocketmqForwarder) serve() { + defer f.wg.Done() + for { + localConn, err := f.listener.Accept() + if err != nil { + f.mu.Lock() + closed := f.closed + f.mu.Unlock() + if !closed { + logger.Warnf("接受 RocketMQ 隧道连接失败:%v", err) + } + return + } + if !f.track(localConn) { + _ = localConn.Close() + return + } + f.wg.Add(1) + go f.handle(localConn) + } +} + +func (f *rocketmqForwarder) handle(localConn net.Conn) { + defer f.wg.Done() + defer f.untrack(localConn) + defer localConn.Close() + + remoteConn, err := f.dialContext(f.ctx, "tcp", f.remoteAddr) + if err != nil { + logger.Warnf("连接 RocketMQ 隧道远端失败:远端=%s 错误=%v", f.remoteAddr, err) + return + } + if !f.track(remoteConn) { + _ = remoteConn.Close() + return + } + defer f.untrack(remoteConn) + defer remoteConn.Close() + + errCh := make(chan error, 2) + go func() { + _, copyErr := io.Copy(remoteConn, localConn) + errCh <- copyErr + }() + go func() { + var copyErr error + if f.rewrite == nil { + _, copyErr = io.Copy(localConn, remoteConn) + } else { + copyErr = relayRocketMQResponses(localConn, remoteConn, f.rewrite) + } + errCh <- copyErr + }() + firstErr := <-errCh + _ = localConn.Close() + _ = remoteConn.Close() + secondErr := <-errCh + if err := rocketmqUnexpectedTunnelError(firstErr, secondErr, f.ctx.Err()); err != nil { + logger.Warnf("转发 RocketMQ 隧道数据失败:远端=%s 错误=%v", f.remoteAddr, err) + } +} + +func rocketmqUnexpectedTunnelError(errs ...error) error { + for _, err := range errs { + if err == nil || errors.Is(err, io.EOF) || errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) { + continue + } + if strings.Contains(strings.ToLower(err.Error()), "use of closed network connection") { + continue + } + return err + } + return nil +} + +func (f *rocketmqForwarder) track(conn net.Conn) bool { + f.mu.Lock() + defer f.mu.Unlock() + if f.closed { + return false + } + f.active[conn] = struct{}{} + return true +} + +func (f *rocketmqForwarder) untrack(conn net.Conn) { + f.mu.Lock() + delete(f.active, conn) + f.mu.Unlock() +} + +func (f *rocketmqForwarder) HasActiveConnections() bool { + f.mu.Lock() + defer f.mu.Unlock() + return len(f.active) > 0 +} + +func (f *rocketmqForwarder) Close() error { + if f == nil { + return nil + } + var closeErr error + f.closeOnce.Do(func() { + f.cancel() + f.mu.Lock() + f.closed = true + connections := make([]net.Conn, 0, len(f.active)) + for conn := range f.active { + connections = append(connections, conn) + } + f.mu.Unlock() + closeErr = f.listener.Close() + for _, conn := range connections { + _ = conn.Close() + } + f.wg.Wait() + }) + return closeErr +} + +func relayRocketMQResponses(dst io.Writer, src io.Reader, rewrite func([]byte) ([]byte, error)) error { + for { + var sizeBuffer [4]byte + if _, err := io.ReadFull(src, sizeBuffer[:]); err != nil { + return err + } + frameSize := int(binary.BigEndian.Uint32(sizeBuffer[:])) + if frameSize < 4 || frameSize > maxRocketMQFrameSize { + return fmt.Errorf("RocketMQ 响应帧长度无效:%d", frameSize) + } + frame := make([]byte, frameSize+4) + copy(frame[:4], sizeBuffer[:]) + if _, err := io.ReadFull(src, frame[4:]); err != nil { + return err + } + rewritten, err := rewrite(frame) + if err != nil { + return err + } + if err := writeAll(dst, rewritten); err != nil { + return err + } + } +} + +func writeAll(dst io.Writer, payload []byte) error { + for len(payload) > 0 { + written, err := dst.Write(payload) + if err != nil { + return err + } + if written <= 0 { + return io.ErrShortWrite + } + payload = payload[written:] + } + return nil +} + +func rewriteRocketMQRouteFrame(frame []byte, mapAddress func(string) (string, error)) ([]byte, error) { + if len(frame) < 8 { + return nil, fmt.Errorf("RocketMQ 响应帧过短") + } + frameSize := int(binary.BigEndian.Uint32(frame[:4])) + if frameSize != len(frame)-4 || frameSize < 4 || frameSize > maxRocketMQFrameSize { + return nil, fmt.Errorf("RocketMQ 响应帧长度无效:%d", frameSize) + } + headerLength := int(binary.BigEndian.Uint32(frame[4:8]) & 0x00ffffff) + bodyOffset := 8 + headerLength + if bodyOffset > len(frame) { + return nil, fmt.Errorf("RocketMQ 响应头长度无效:%d", headerLength) + } + body, changed, err := rewriteRocketMQRouteBody(frame[bodyOffset:], mapAddress) + if err != nil || !changed { + return frame, err + } + + newFrameSize := 4 + headerLength + len(body) + rewritten := make([]byte, newFrameSize+4) + binary.BigEndian.PutUint32(rewritten[:4], uint32(newFrameSize)) + copy(rewritten[4:bodyOffset], frame[4:bodyOffset]) + copy(rewritten[bodyOffset:], body) + return rewritten, nil +} + +func rewriteRocketMQRouteBody(body []byte, mapAddress func(string) (string, error)) ([]byte, bool, error) { + var payload map[string]json.RawMessage + if len(body) == 0 || json.Unmarshal(body, &payload) != nil { + return body, false, nil + } + brokerDataJSON, ok := payload["brokerDatas"] + if !ok { + return body, false, nil + } + var brokerDatas []map[string]json.RawMessage + if err := json.Unmarshal(brokerDataJSON, &brokerDatas); err != nil { + return nil, false, fmt.Errorf("解析 RocketMQ Broker 路由失败:%w", err) + } + + changed := false + for _, brokerData := range brokerDatas { + addressesJSON, exists := brokerData["brokerAddrs"] + if !exists { + continue + } + var addresses map[string]string + if err := json.Unmarshal(addressesJSON, &addresses); err != nil { + return nil, false, fmt.Errorf("解析 RocketMQ Broker 地址失败:%w", err) + } + brokerChanged := false + for id, address := range addresses { + forwarded, err := mapAddress(address) + if err != nil { + return nil, false, err + } + if forwarded != address { + addresses[id] = forwarded + brokerChanged = true + changed = true + } + } + if brokerChanged { + rewritten, err := json.Marshal(addresses) + if err != nil { + return nil, false, err + } + brokerData["brokerAddrs"] = rewritten + } + } + if !changed { + return body, false, nil + } + rewrittenBrokerDatas, err := json.Marshal(brokerDatas) + if err != nil { + return nil, false, err + } + payload["brokerDatas"] = rewrittenBrokerDatas + rewrittenBody, err := json.Marshal(payload) + if err != nil { + return nil, false, err + } + return rewrittenBody, true, nil +} diff --git a/internal/db/rocketmq_tunnel_test.go b/internal/db/rocketmq_tunnel_test.go new file mode 100644 index 00000000..57f6cad7 --- /dev/null +++ b/internal/db/rocketmq_tunnel_test.go @@ -0,0 +1,500 @@ +package db + +import ( + "bufio" + "bytes" + "context" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "reflect" + "strconv" + "strings" + "sync" + "testing" + "time" + + "GoNavi-Wails/internal/connection" +) + +func TestRewriteRocketMQRouteFrameMapsBrokerAddresses(t *testing.T) { + body := []byte(`{"brokerDatas":[{"brokerName":"broker-a","brokerAddrs":{"0":"10.0.0.10:10911","1":"10.0.0.11:10911"}}],"queueDatas":[]}`) + frame := rocketmqTestFrame([]byte(`{"code":0}`), body) + mapped := map[string]string{ + "10.0.0.10:10911": "127.0.0.1:31001", + "10.0.0.11:10911": "127.0.0.1:31002", + } + calls := make([]string, 0, len(mapped)) + + rewritten, err := rewriteRocketMQRouteFrame(frame, func(address string) (string, error) { + calls = append(calls, address) + return mapped[address], nil + }) + if err != nil { + t.Fatalf("rewriteRocketMQRouteFrame: %v", err) + } + if got := int(binary.BigEndian.Uint32(rewritten[:4])); got != len(rewritten)-4 { + t.Fatalf("frame size = %d, want %d", got, len(rewritten)-4) + } + if got := int(binary.BigEndian.Uint32(rewritten[4:8]) & 0x00ffffff); got != len(`{"code":0}`) { + t.Fatalf("header size = %d, want %d", got, len(`{"code":0}`)) + } + + bodyOffset := 8 + len(`{"code":0}`) + var payload struct { + BrokerDatas []struct { + BrokerAddrs map[string]string `json:"brokerAddrs"` + } `json:"brokerDatas"` + } + if err := json.Unmarshal(rewritten[bodyOffset:], &payload); err != nil { + t.Fatalf("decode rewritten route body: %v", err) + } + wantAddresses := map[string]string{"0": "127.0.0.1:31001", "1": "127.0.0.1:31002"} + if len(payload.BrokerDatas) != 1 || !reflect.DeepEqual(payload.BrokerDatas[0].BrokerAddrs, wantAddresses) { + t.Fatalf("rewritten broker addresses = %#v, want %#v", payload.BrokerDatas, wantAddresses) + } + if len(calls) != 2 { + t.Fatalf("broker mapper calls = %v, want both addresses", calls) + } +} + +func TestRewriteRocketMQRouteFrameKeepsNonRouteResponse(t *testing.T) { + frame := rocketmqTestFrame([]byte(`{"code":0}`), []byte(`{"topicList":["orders.events"]}`)) + rewritten, err := rewriteRocketMQRouteFrame(frame, func(address string) (string, error) { + t.Fatalf("broker mapper called for non-route response: %s", address) + return "", nil + }) + if err != nil { + t.Fatalf("rewriteRocketMQRouteFrame: %v", err) + } + if !reflect.DeepEqual(rewritten, frame) { + t.Fatalf("non-route response changed: got=%q want=%q", rewritten, frame) + } +} + +func TestNormalizeRocketMQTunnelConfigDefaultsSSHPort(t *testing.T) { + config, err := normalizeRocketMQTunnelConfig(connection.ConnectionConfig{ + UseSSH: true, + SSH: connection.SSHConfig{ + Host: "ssh.internal.test", + User: "ssh-user", + }, + }) + if err != nil { + t.Fatalf("normalizeRocketMQTunnelConfig: %v", err) + } + if config.SSH.Port != 22 { + t.Fatalf("SSH port = %d, want 22", config.SSH.Port) + } +} + +func TestRelayRocketMQResponsesHandlesShortWrites(t *testing.T) { + frame := rocketmqTestFrame([]byte(`{"code":0}`), []byte(`{"topicList":["orders.events"]}`)) + writer := &rocketmqShortWriter{max: 3} + err := relayRocketMQResponses(writer, bytes.NewReader(frame), func(payload []byte) ([]byte, error) { + return payload, nil + }) + if !errors.Is(err, io.EOF) { + t.Fatalf("relayRocketMQResponses error = %v, want EOF", err) + } + if !bytes.Equal(writer.payload, frame) { + t.Fatalf("relayed frame = %q, want %q", writer.payload, frame) + } +} + +func TestRocketMQForwarderCloseCancelsInFlightDial(t *testing.T) { + dialStarted := make(chan struct{}) + dialCanceled := make(chan struct{}) + forwarder, err := newRocketMQForwarder("nameserver.internal.test:9876", func(ctx context.Context, network, address string) (net.Conn, error) { + close(dialStarted) + <-ctx.Done() + close(dialCanceled) + return nil, ctx.Err() + }, nil) + if err != nil { + t.Fatalf("newRocketMQForwarder: %v", err) + } + + client, err := net.DialTimeout("tcp", forwarder.LocalAddr(), time.Second) + if err != nil { + t.Fatalf("dial local forwarder: %v", err) + } + defer client.Close() + select { + case <-dialStarted: + case <-time.After(time.Second): + t.Fatal("forwarder dial did not start") + } + + closed := make(chan error, 1) + go func() { closed <- forwarder.Close() }() + select { + case err := <-closed: + if err != nil { + t.Fatalf("Close: %v", err) + } + case <-time.After(time.Second): + t.Fatal("Close remained blocked while dial was in flight") + } + select { + case <-dialCanceled: + case <-time.After(time.Second): + t.Fatal("in-flight dial did not observe cancellation") + } +} + +func TestRocketMQTunnelEvictsIdleBrokerForwarders(t *testing.T) { + originalNow := rocketMQTunnelNow + now := time.Date(2026, 8, 9, 12, 0, 0, 0, time.UTC) + rocketMQTunnelNow = func() time.Time { return now } + t.Cleanup(func() { rocketMQTunnelNow = originalNow }) + + tunnels := newRocketMQTunnelSet((&net.Dialer{}).DialContext) + t.Cleanup(func() { _ = tunnels.Close() }) + firstLocal, err := tunnels.forwardBrokerAddress("127.0.0.1:10911") + if err != nil { + t.Fatalf("forward first broker: %v", err) + } + now = now.Add(5 * time.Minute) + staleLocal, err := tunnels.forwardBrokerAddress("127.0.0.1:10912") + if err != nil { + t.Fatalf("forward stale broker: %v", err) + } + now = now.Add(11 * time.Minute) + firstLocalAgain, err := tunnels.forwardBrokerAddress("127.0.0.1:10911") + if err != nil { + t.Fatalf("refresh first broker: %v", err) + } + if firstLocalAgain != firstLocal { + t.Fatalf("active broker forwarder changed: got=%s want=%s", firstLocalAgain, firstLocal) + } + if conn, dialErr := net.DialTimeout("tcp", staleLocal, 50*time.Millisecond); dialErr == nil { + _ = conn.Close() + t.Fatalf("idle broker forwarder still accepts connections: %s", staleLocal) + } +} + +func TestRocketMQTunnelRewritesFragmentedNameServerRoutesAndForwardsBroker(t *testing.T) { + broker := startRocketMQEchoServer(t, "broker-ping", "broker-pong") + nameserver := startRocketMQNameServerStub(t, broker.Addr().String()) + proxy := startRocketMQHTTPConnectProxy(t) + + proxyHost, proxyPort := rocketmqTestHostPort(t, proxy.Addr().String()) + nameserverHost, nameserverPort := rocketmqTestHostPort(t, nameserver.Addr().String()) + config, tunnels, err := prepareRocketMQTunnel(connection.ConnectionConfig{ + Type: "rocketmq", + Host: nameserverHost, + Port: nameserverPort, + UseProxy: true, + Proxy: connection.ProxyConfig{ + Type: "http", + Host: proxyHost, + Port: proxyPort, + }, + }) + if err != nil { + t.Fatalf("prepareRocketMQTunnel: %v", err) + } + t.Cleanup(func() { _ = tunnels.Close() }) + + exerciseRocketMQForwardedRoute(t, config, broker.Addr().String()) + targets := proxy.Targets() + if !containsString(targets, nameserver.Addr().String()) || !containsString(targets, broker.Addr().String()) { + t.Fatalf("HTTP CONNECT targets = %v, want NameServer and Broker", targets) + } +} + +func TestRocketMQSSHTunnelForwardsNameServerAndBroker(t *testing.T) { + broker := startRocketMQEchoServer(t, "broker-ping", "broker-pong") + nameserver := startRocketMQNameServerStub(t, broker.Addr().String()) + originalDial := rocketMQDialContextThroughSSH + var mu sync.Mutex + var targets []string + var sshConfigs []connection.SSHConfig + rocketMQDialContextThroughSSH = func(ctx context.Context, config connection.SSHConfig, network, address string) (net.Conn, error) { + mu.Lock() + targets = append(targets, address) + sshConfigs = append(sshConfigs, config) + mu.Unlock() + return (&net.Dialer{}).DialContext(ctx, network, address) + } + t.Cleanup(func() { rocketMQDialContextThroughSSH = originalDial }) + + nameserverHost, nameserverPort := rocketmqTestHostPort(t, nameserver.Addr().String()) + config, tunnels, err := prepareRocketMQTunnel(connection.ConnectionConfig{ + Type: "rocketmq", + Host: nameserverHost, + Port: nameserverPort, + UseSSH: true, + SSH: connection.SSHConfig{ + Host: "ssh.internal.test", + User: "ssh-user", + }, + }) + if err != nil { + t.Fatalf("prepareRocketMQTunnel: %v", err) + } + t.Cleanup(func() { _ = tunnels.Close() }) + + exerciseRocketMQForwardedRoute(t, config, broker.Addr().String()) + mu.Lock() + defer mu.Unlock() + if !containsString(targets, nameserver.Addr().String()) || !containsString(targets, broker.Addr().String()) { + t.Fatalf("SSH targets = %v, want NameServer and Broker", targets) + } + for _, sshConfig := range sshConfigs { + if sshConfig.Port != 22 { + t.Fatalf("SSH port = %d, want normalized port 22", sshConfig.Port) + } + } +} + +func exerciseRocketMQForwardedRoute(t *testing.T, config connection.ConnectionConfig, originalBroker string) { + t.Helper() + nameServerConn, err := net.DialTimeout("tcp", rocketmqFormatHostPort(config.Host, config.Port), time.Second) + if err != nil { + t.Fatalf("dial forwarded NameServer: %v", err) + } + if err := writeAll(nameServerConn, rocketmqTestFrame([]byte(`{"code":105}`), nil)); err != nil { + t.Fatalf("write NameServer request: %v", err) + } + firstFrame, err := readRocketMQTestFrame(nameServerConn) + if err != nil { + t.Fatalf("read first NameServer response: %v", err) + } + if !strings.Contains(string(rocketmqTestFrameBody(t, firstFrame)), "topicList") { + t.Fatalf("first response was not preserved: %q", firstFrame) + } + secondFrame, err := readRocketMQTestFrame(nameServerConn) + if err != nil { + t.Fatalf("read route response: %v", err) + } + _ = nameServerConn.Close() + + var route struct { + BrokerDatas []struct { + BrokerAddrs map[string]string `json:"brokerAddrs"` + } `json:"brokerDatas"` + } + if err := json.Unmarshal(rocketmqTestFrameBody(t, secondFrame), &route); err != nil { + t.Fatalf("decode rewritten route: %v", err) + } + if len(route.BrokerDatas) != 1 { + t.Fatalf("rewritten route brokers = %#v", route.BrokerDatas) + } + forwardedBroker := route.BrokerDatas[0].BrokerAddrs["0"] + if forwardedBroker == "" || forwardedBroker == originalBroker { + t.Fatalf("broker address was not rewritten: %q", forwardedBroker) + } + brokerConn, err := net.DialTimeout("tcp", forwardedBroker, time.Second) + if err != nil { + t.Fatalf("dial forwarded Broker: %v", err) + } + if err := writeAll(brokerConn, []byte("broker-ping")); err != nil { + t.Fatalf("write Broker payload: %v", err) + } + response := make([]byte, len("broker-pong")) + if _, err := io.ReadFull(brokerConn, response); err != nil { + t.Fatalf("read Broker response: %v", err) + } + _ = brokerConn.Close() + if string(response) != "broker-pong" { + t.Fatalf("Broker response = %q, want broker-pong", response) + } +} + +func rocketmqTestFrame(header []byte, body []byte) []byte { + frameSize := 4 + len(header) + len(body) + frame := make([]byte, frameSize+4) + binary.BigEndian.PutUint32(frame[:4], uint32(frameSize)) + binary.BigEndian.PutUint32(frame[4:8], uint32(len(header))) + copy(frame[8:], header) + copy(frame[8+len(header):], body) + return frame +} + +func readRocketMQTestFrame(reader io.Reader) ([]byte, error) { + var sizeBuffer [4]byte + if _, err := io.ReadFull(reader, sizeBuffer[:]); err != nil { + return nil, err + } + frameSize := int(binary.BigEndian.Uint32(sizeBuffer[:])) + frame := make([]byte, frameSize+4) + copy(frame[:4], sizeBuffer[:]) + _, err := io.ReadFull(reader, frame[4:]) + return frame, err +} + +func rocketmqTestFrameBody(t *testing.T, frame []byte) []byte { + t.Helper() + if len(frame) < 8 { + t.Fatalf("RocketMQ test frame too short: %d", len(frame)) + } + headerLength := int(binary.BigEndian.Uint32(frame[4:8]) & 0x00ffffff) + bodyOffset := 8 + headerLength + if bodyOffset > len(frame) { + t.Fatalf("RocketMQ test frame header length = %d, frame length = %d", headerLength, len(frame)) + } + return frame[bodyOffset:] +} + +func rocketmqTestHostPort(t *testing.T, address string) (string, int) { + t.Helper() + host, portText, err := net.SplitHostPort(address) + if err != nil { + t.Fatalf("split address %q: %v", address, err) + } + port, err := strconv.Atoi(portText) + if err != nil { + t.Fatalf("parse port %q: %v", portText, err) + } + return host, port +} + +func startRocketMQEchoServer(t *testing.T, request string, response string) net.Listener { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen for Broker stub: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + defer conn.Close() + payload := make([]byte, len(request)) + if _, readErr := io.ReadFull(conn, payload); readErr != nil || string(payload) != request { + return + } + _, _ = io.WriteString(conn, response) + }() + return listener +} + +func startRocketMQNameServerStub(t *testing.T, brokerAddress string) net.Listener { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen for NameServer stub: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + defer conn.Close() + if _, readErr := readRocketMQTestFrame(conn); readErr != nil { + return + } + frames := append( + rocketmqTestFrame([]byte(`{"code":0}`), []byte(`{"topicList":["orders.events"]}`)), + rocketmqTestFrame( + []byte(`{"code":0}`), + []byte(fmt.Sprintf(`{"brokerDatas":[{"brokerName":"broker-a","brokerAddrs":{"0":%q}}],"queueDatas":[]}`, brokerAddress)), + )..., + ) + for _, value := range frames { + if _, writeErr := conn.Write([]byte{value}); writeErr != nil { + return + } + } + }() + return listener +} + +type rocketmqShortWriter struct { + max int + payload []byte +} + +func (w *rocketmqShortWriter) Write(payload []byte) (int, error) { + written := len(payload) + if written > w.max { + written = w.max + } + w.payload = append(w.payload, payload[:written]...) + return written, nil +} + +type rocketmqHTTPConnectProxy struct { + listener net.Listener + mu sync.Mutex + targets []string +} + +func startRocketMQHTTPConnectProxy(t *testing.T) *rocketmqHTTPConnectProxy { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen for HTTP CONNECT proxy: %v", err) + } + proxy := &rocketmqHTTPConnectProxy{listener: listener} + t.Cleanup(func() { _ = listener.Close() }) + go proxy.serve() + return proxy +} + +func (p *rocketmqHTTPConnectProxy) Addr() net.Addr { + return p.listener.Addr() +} + +func (p *rocketmqHTTPConnectProxy) Targets() []string { + p.mu.Lock() + defer p.mu.Unlock() + return append([]string(nil), p.targets...) +} + +func (p *rocketmqHTTPConnectProxy) serve() { + for { + conn, err := p.listener.Accept() + if err != nil { + return + } + go p.handle(conn) + } +} + +func (p *rocketmqHTTPConnectProxy) handle(client net.Conn) { + defer client.Close() + request, err := http.ReadRequest(bufio.NewReader(client)) + if err != nil { + return + } + _ = request.Body.Close() + target := request.Host + if request.Method != http.MethodConnect || target == "" { + return + } + p.mu.Lock() + p.targets = append(p.targets, target) + p.mu.Unlock() + + upstream, err := (&net.Dialer{}).DialContext(context.Background(), "tcp", target) + if err != nil { + return + } + defer upstream.Close() + if _, err := io.WriteString(client, "HTTP/1.1 200 Connection Established\r\n\r\n"); err != nil { + return + } + errCh := make(chan error, 2) + go func() { + _, copyErr := io.Copy(upstream, client) + errCh <- copyErr + }() + go func() { + _, copyErr := io.Copy(client, upstream) + errCh <- copyErr + }() + <-errCh + _ = client.Close() + _ = upstream.Close() + <-errCh +}