🐛 fix(rocketmq): 支持经隧道访问隔离网络 (#886)

Fixes #886
This commit is contained in:
mango
2026-08-09 14:20:33 +08:00
parent bf5a1b7681
commit f887deb44c
6 changed files with 1205 additions and 5 deletions

View File

@@ -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

View File

@@ -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)
}
})
}
}

View File

@@ -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 = ""

View File

@@ -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{

View File

@@ -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
}

View File

@@ -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
}