mirror of
https://github.com/Syngnat/GoNavi.git
synced 2026-08-15 03:13:39 +08:00
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 = ""
|
||||
|
||||
@@ -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{
|
||||
|
||||
534
internal/db/rocketmq_tunnel.go
Normal file
534
internal/db/rocketmq_tunnel.go
Normal 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
|
||||
}
|
||||
500
internal/db/rocketmq_tunnel_test.go
Normal file
500
internal/db/rocketmq_tunnel_test.go
Normal 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
|
||||
}
|
||||
Reference in New Issue
Block a user