diff --git a/.github/scripts/bump-agent-versions.mjs b/.github/scripts/bump-agent-versions.mjs index e58f3d572..48676b909 100644 --- a/.github/scripts/bump-agent-versions.mjs +++ b/.github/scripts/bump-agent-versions.mjs @@ -46,6 +46,7 @@ const nativeDriverDirectories = { duckdb: "duckdb", oracle: "oracle-go", kingbase: "kingbase-go", + rabbitmq: "rabbitmq", }; function resolveAgentModule(moduleName, { legacyStandaloneModules, moduleExists, readModuleFile }) { diff --git a/.github/scripts/bump-agent-versions.test.mjs b/.github/scripts/bump-agent-versions.test.mjs index 354c2925f..c3bd0f94a 100644 --- a/.github/scripts/bump-agent-versions.test.mjs +++ b/.github/scripts/bump-agent-versions.test.mjs @@ -29,3 +29,14 @@ test("bumps DuckDB after its initial release", () => { assert.equal(result.versions.duckdb, "0.1.1"); }); + +test("bumps the native RabbitMQ agent from its Go directory", () => { + const result = evaluateAgentVersionBump({ + versions: { rabbitmq: "0.1.0" }, + changedFiles: ["agents/drivers/rabbitmq/main.go"], + moduleExists: (path) => path === "agents/drivers/rabbitmq", + readModuleFile: () => "", + }); + + assert.equal(result.versions.rabbitmq, "0.1.1"); +}); diff --git a/.github/workflows/agents-release.yml b/.github/workflows/agents-release.yml index 3c06978ca..45aef7246 100644 --- a/.github/workflows/agents-release.yml +++ b/.github/workflows/agents-release.yml @@ -217,6 +217,45 @@ jobs: name: xugu-native path: "release-native/dbx-agent-xugu-*" + build-rabbitmq-native: + needs: [bump-versions] + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.22.x" + - name: Test RabbitMQ native agent + working-directory: agents/drivers/rabbitmq + run: go test ./... + - name: Cross-compile RabbitMQ native agent + shell: bash + run: | + mkdir -p release-native + cd agents/drivers/rabbitmq + declare -A TARGETS=( + ["macos-aarch64"]="darwin/arm64" + ["macos-x64"]="darwin/amd64" + ["linux-aarch64"]="linux/arm64" + ["linux-x64"]="linux/amd64" + ["windows-aarch64"]="windows/arm64" + ["windows-x64"]="windows/amd64" + ) + for platform in "${!TARGETS[@]}"; do + IFS=/ read -r goos goarch <<< "${TARGETS[$platform]}" + output="../../../release-native/dbx-agent-rabbitmq-${platform}" + if [[ "$goos" == "windows" ]]; then + output="${output}.exe" + fi + echo "Building $platform ($goos/$goarch)" + CGO_ENABLED=0 GOOS="$goos" GOARCH="$goarch" go build -trimpath -ldflags="-s -w" -o "$output" . + done + ls -lh ../../../release-native + - uses: actions/upload-artifact@v4 + with: + name: rabbitmq-native + path: "release-native/dbx-agent-rabbitmq-*" + build-kingbase-native: needs: [bump-versions] runs-on: ubuntu-latest @@ -436,7 +475,7 @@ jobs: path: "dbx-jre-*.tar.zst" release: - needs: [bump-versions, commit-versions, build-agents, build-oracle-native, build-xugu-native, build-kingbase-native, build-duckdb-native, build-jre] + needs: [bump-versions, commit-versions, build-agents, build-oracle-native, build-xugu-native, build-rabbitmq-native, build-kingbase-native, build-duckdb-native, build-jre] runs-on: ubuntu-latest steps: - name: Create DBX bot release token @@ -466,6 +505,7 @@ jobs: find artifacts/agent-jars -name '*.jar' -exec cp {} release/ \; find artifacts/oracle-native -type f -name 'dbx-agent-oracle-*' -exec cp {} release/ \; find artifacts/xugu-native -type f -name 'dbx-agent-xugu-*' -exec cp {} release/ \; + find artifacts/rabbitmq-native -type f -name 'dbx-agent-rabbitmq-*' -exec cp {} release/ \; find artifacts/kingbase-native -type f -name 'dbx-agent-kingbase-*' -exec cp {} release/ \; find artifacts/duckdb-native-* -type f -name 'dbx-agent-duckdb-*' -exec cp {} release/ \; find artifacts -name 'dbx-jre-*.tar.zst' -exec cp {} release/ \; @@ -540,6 +580,7 @@ jobs: kingbase) echo "人大金仓 KingbaseES" ;; duckdb) echo "DuckDB" ;; xugu) echo "虚谷 XuguDB" ;; + rabbitmq) echo "RabbitMQ" ;; *) echo "$name" ;; esac } @@ -603,7 +644,7 @@ jobs: [ -n "$DRIVERS" ] && DRIVERS="${DRIVERS},"$'\n' DRIVERS="${DRIVERS}$(generate_jar_entry "$name" "$label" "$f" "$jre_key" "$version" "$external_driver" "$native_json")" done - for name in oracle xugu kingbase duckdb; do + for name in oracle xugu kingbase duckdb rabbitmq; do version=$(get_module_version "$name") [ -f "release/dbx-agent-${name}-${version}.jar" ] && continue native_json=$(generate_native_platforms "$name" "$version") @@ -673,6 +714,7 @@ jobs: duckdb) echo "DuckDB" ;; oracle) echo "Oracle" ;; xugu) echo "虚谷 XuguDB" ;; + rabbitmq) echo "RabbitMQ" ;; *) echo "$name" ;; esac } diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b5085679c..40ccb99cb 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -513,6 +513,10 @@ jobs: run: GONOSUMDB=gitee.com/XuguDB/go-xugu-driver go test ./... working-directory: agents/drivers/xugu + - name: RabbitMQ native agent tests + run: go test ./... + working-directory: agents/drivers/rabbitmq + - name: Oracle native agent build run: CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -trimpath -ldflags="-s -w" -o /tmp/dbx-agent-oracle-linux-x64 . working-directory: agents/drivers/oracle-go @@ -521,6 +525,46 @@ jobs: run: GONOSUMDB=gitee.com/XuguDB/go-xugu-driver CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -trimpath -ldflags="-s -w" -o /tmp/dbx-agent-xugu-linux-x64 . working-directory: agents/drivers/xugu + - name: RabbitMQ native agent build + run: CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -trimpath -ldflags="-s -w" -o /tmp/dbx-agent-rabbitmq-linux-x64 . + working-directory: agents/drivers/rabbitmq + + - name: RabbitMQ native agent integration tests + shell: bash + working-directory: agents/drivers/rabbitmq + run: | + set -euo pipefail + for version in 3.13 4.3; do + name="dbx-rabbitmq-${version//./-}" + docker run -d --name "$name" \ + -e RABBITMQ_DEFAULT_USER=dbx \ + -e RABBITMQ_DEFAULT_PASS=dbx-password \ + -p 5672:5672 -p 15672:15672 \ + "rabbitmq:${version}-management" + cleanup() { + docker rm -f "$name" >/dev/null 2>&1 || true + } + trap cleanup EXIT + ready=false + for _ in $(seq 1 60); do + if docker exec "$name" rabbitmq-diagnostics -q ping >/dev/null 2>&1; then + ready=true + break + fi + sleep 2 + done + if [ "$ready" != "true" ]; then + docker logs "$name" + exit 1 + fi + RABBITMQ_INTEGRATION=1 \ + RABBITMQ_USERNAME=dbx \ + RABBITMQ_PASSWORD=dbx-password \ + go test -run '^TestRabbitMQIntegration$' -count=1 ./... + cleanup + trap - EXIT + done + - name: Java agent tests and packages run: ./gradlew test shadowJar --continue diff --git a/agents/README.md b/agents/README.md index 3ccf969f1..6f24213c3 100644 --- a/agents/README.md +++ b/agents/README.md @@ -44,12 +44,12 @@ Each agent runs as a standalone process and communicates with DBX via stdin/stdo | iotdb | Apache IoTDB | IoTDB JDBC | | etcd | etcd | jetcd | | zookeeper | Apache ZooKeeper | Apache Curator | -| rabbitmq | RabbitMQ | RabbitMQ AMQP Java client | +| rabbitmq | RabbitMQ | amqp091-go native agent | ## Multi-JRE Support -Most Java agents target JRE 21. Native agents, such as `duckdb`, `oracle`, `kingbase`, and `xugu`, do not require a JRE. DBX downloads and manages the JRE 21 installation automatically for Java agents. +Most Java agents target JRE 21. Native agents, such as `duckdb`, `oracle`, `kingbase`, `xugu`, and `rabbitmq`, do not require a JRE. DBX downloads and manages the JRE 21 installation automatically for Java agents. ## JDBC Connection Pooling @@ -75,7 +75,7 @@ Set `DBX_AGENT_JDBC_POOL_ENABLED=false` for a runtime-level compatibility fallba For new agents, prefer a **native (Go or Rust) driver** over a Java/JDBC agent whenever a mature, license-compatible native driver is available. Native agents ship as a single self-contained executable with no JRE, which significantly reduces memory footprint and startup time — the JVM baseline that every Java agent pays even when idle is avoided entirely. -- **Native (C++/Go/Rust)** — preferred when a usable native driver exists. See `drivers/duckdb`, `drivers/oracle-go` (go-ora), `drivers/kingbase-go` (gokb), and `drivers/xugu` as reference implementations. No JRE download or management is needed. +- **Native (C++/Go/Rust)** — preferred when a usable native driver exists. See `drivers/duckdb`, `drivers/oracle-go` (go-ora), `drivers/kingbase-go` (gokb), `drivers/xugu`, and `drivers/rabbitmq` (amqp091-go) as reference implementations. No JRE download or management is needed. - **Java/JDBC** — the default fallback when only a JDBC driver exists for the database, or when the native driver is immature or unmaintained. Most agents still fall in this category. Native agents implement the same JSON-RPC contract and `versions.json` registration as Java agents; they ship an `agent` executable instead of `agent.jar`. If both native and Java source implementations exist for the same database, publish only the native artifact unless the Java variant has a separately registered compatibility profile, such as `oracle-legacy` / `oracle-10g`. @@ -89,9 +89,10 @@ Requires JDK 21 (Gradle toolchain auto-downloads if needed). (cd drivers/oracle-go && go build -o agent .) (cd drivers/kingbase-go && go build -o agent .) (cd drivers/xugu && go build -o agent .) +(cd drivers/rabbitmq && go build -o agent .) ``` -Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/oracle-go`, `drivers/kingbase-go`, and `drivers/xugu`. +Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/oracle-go`, `drivers/kingbase-go`, `drivers/xugu`, and `drivers/rabbitmq`. ### Local DBX Runtime Test @@ -105,7 +106,7 @@ cp agents/drivers//build/libs/*-all.jar ~/.dbx/agents/drivers/ Restart DBX or disconnect and reconnect the database so the new agent process loads the replacement JAR. -Native agents such as `oracle`, `kingbase`, and `xugu` use the `agent` executable in the driver directory instead of `agent.jar`. +Native agents such as `oracle`, `kingbase`, `xugu`, and `rabbitmq` use the `agent` executable in the driver directory instead of `agent.jar`. ## Versioning diff --git a/agents/README.zh-CN.md b/agents/README.zh-CN.md index 5c47c95a5..7e5b8d5a4 100644 --- a/agents/README.zh-CN.md +++ b/agents/README.zh-CN.md @@ -44,12 +44,12 @@ DBX 的 Agent 驱动 —— 通过 JDBC 和原生数据库驱动支持各种数 | iotdb | Apache IoTDB | IoTDB JDBC | | etcd | etcd | jetcd | | zookeeper | Apache ZooKeeper | Apache Curator | -| rabbitmq | RabbitMQ | RabbitMQ AMQP Java client | +| rabbitmq | RabbitMQ | amqp091-go 原生 agent | ## 多 JRE 支持 -多数 Java agent 以 JRE 21 为目标。原生 agent(如 `oracle`、`kingbase` 和 `xugu`)不需要 JRE。对 Java agent,DBX 会自动下载并管理 JRE 21 安装。 +多数 Java agent 以 JRE 21 为目标。原生 agent(如 `oracle`、`kingbase`、`xugu` 和 `rabbitmq`)不需要 JRE。对 Java agent,DBX 会自动下载并管理 JRE 21 安装。 ## JDBC 连接池 @@ -75,7 +75,7 @@ HikariCP 会直接打进启用连接池的 Agent JAR。已经使用 DBX 托管 J 对于新 agent,只要存在成熟、许可证兼容的原生驱动,优先选择**原生(Go 或 Rust)驱动**而非 Java/JDBC agent。原生 agent 以单一自包含可执行文件发布,无需 JRE,可显著降低内存占用和启动时间 —— 完全避开 Java agent 即便空闲也要付出的 JVM 基线开销。 -- **原生(Go/Rust)** —— 存在可用原生驱动时首选。参考 `drivers/oracle-go`(go-ora)、`drivers/kingbase-go`(gokb)和 `drivers/xugu`。无需 JRE 下载与管理。 +- **原生(Go/Rust)** —— 存在可用原生驱动时首选。参考 `drivers/oracle-go`(go-ora)、`drivers/kingbase-go`(gokb)、`drivers/xugu` 和 `drivers/rabbitmq`(amqp091-go)。无需 JRE 下载与管理。 - **Java/JDBC** —— 当某数据库只有 JDBC 驱动,或原生驱动不成熟、缺乏维护时的默认兜底方案。多数 agent 仍属此类。 原生 agent 实现与 Java agent 相同的 JSON-RPC 契约和 `versions.json` 登记;它发布的是 `agent` 可执行文件而非 `agent.jar`。若同一数据库同时保留原生和 Java 源码实现,默认只发布原生产物;只有 Java 变体以独立兼容配置登记时才同时发布,例如 `oracle-legacy` / `oracle-10g`。 @@ -89,9 +89,10 @@ HikariCP 会直接打进启用连接池的 Agent JAR。已经使用 DBX 托管 J (cd drivers/oracle-go && go build -o agent .) (cd drivers/kingbase-go && go build -o agent .) (cd drivers/xugu && go build -o agent .) +(cd drivers/rabbitmq && go build -o agent .) ``` -产物 JAR 在 `drivers/{module}/build/libs/`。原生 agent 从 `drivers/oracle-go`、`drivers/kingbase-go` 和 `drivers/xugu` 构建。 +产物 JAR 在 `drivers/{module}/build/libs/`。原生 agent 从 `drivers/oracle-go`、`drivers/kingbase-go`、`drivers/xugu` 和 `drivers/rabbitmq` 构建。 ### 本地 DBX 运行时测试 @@ -105,7 +106,7 @@ cp agents/drivers//build/libs/*-all.jar ~/.dbx/agents/drivers/ 重启 DBX 或断开重连数据库,使新 agent 进程加载替换后的 JAR。 -`oracle`、`kingbase` 和 `xugu` 等原生 agent 使用驱动目录下的 `agent` 可执行文件而非 `agent.jar`。 +`oracle`、`kingbase`、`xugu` 和 `rabbitmq` 等原生 agent 使用驱动目录下的 `agent` 可执行文件而非 `agent.jar`。 ## 版本管理 diff --git a/agents/build.gradle b/agents/build.gradle index 8621eb559..181160ece 100644 --- a/agents/build.gradle +++ b/agents/build.gradle @@ -3,7 +3,7 @@ plugins { } def infrastructureProjects = ['common', 'test-support'] as Set -def legacyStandaloneProjects = ['mongodb', 'kafka', 'rocketmq', 'rabbitmq'] as Set +def legacyStandaloneProjects = ['mongodb', 'kafka', 'rocketmq'] as Set def pooledJdbcProjects = [ 'access', 'bigquery', 'cassandra', 'dameng', 'databend', 'databricks', 'db2', 'exasol', 'firebird', 'gbase8a', 'gbase8s', 'goldendb', 'h2', 'h2-legacy', 'highgo', 'hive', diff --git a/agents/drivers/rabbitmq/bench/agent_compare.go b/agents/drivers/rabbitmq/bench/agent_compare.go new file mode 100644 index 000000000..f7774c4ef --- /dev/null +++ b/agents/drivers/rabbitmq/bench/agent_compare.go @@ -0,0 +1,523 @@ +package main + +import ( + "bufio" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "net" + "net/http" + "os" + "os/exec" + "runtime" + "sort" + "strconv" + "strings" + "time" +) + +type agentSpec struct { + Name string + Command []string + ArtifactPath string +} + +type agentProcess struct { + command *exec.Cmd + stdin io.WriteCloser + reader *bufio.Scanner + nextID int64 +} + +type agentResponse struct { + ID int64 `json:"id"` + Result json.RawMessage `json:"result"` + Error *struct { + Message string `json:"message"` + } `json:"error"` +} + +type benchmarkResult struct { + Agent string `json:"agent"` + Workload string `json:"workload"` + Round int `json:"round"` + Operations int `json:"operations"` + Errors int `json:"errors"` + DurationMS float64 `json:"duration_ms"` + QPS float64 `json:"qps"` + MeanMS float64 `json:"mean_ms"` + P50MS float64 `json:"p50_ms"` + P95MS float64 `json:"p95_ms"` + P99MS float64 `json:"p99_ms"` + ReadyRSSKB int64 `json:"ready_rss_kb,omitempty"` + PostLoadRSSKB int64 `json:"post_load_rss_kb,omitempty"` + ArtifactBytes int64 `json:"artifact_bytes,omitempty"` +} + +type benchmarkMetadata struct { + Type string `json:"type"` + GOOS string `json:"goos"` + GOARCH string `json:"goarch"` + Rounds int `json:"rounds"` + StartupWarmups int `json:"startup_warmups"` + StartupIterations int `json:"startup_iterations"` + WarmupRequests int `json:"warmup_requests"` + RPCRequests int `json:"rpc_requests"` + ManagementCalls int `json:"management_requests"` +} + +func main() { + rounds := envInt("BENCH_ROUNDS", 5) + startupWarmups := envInt("BENCH_STARTUP_WARMUPS", 3) + startupIterations := envInt("BENCH_STARTUPS", 30) + warmupRequests := envInt("BENCH_WARMUP_REQUESTS", 500) + rpcRequests := envInt("BENCH_RPC_REQUESTS", 5000) + managementRequests := envInt("BENCH_MANAGEMENT_REQUESTS", 1000) + + agents := []agentSpec{ + { + Name: "java", + Command: javaAgentCommand(requiredEnv("JAVA_AGENT_JAR")), + ArtifactPath: requiredEnv("JAVA_AGENT_JAR"), + }, + { + Name: "go", + Command: []string{requiredEnv("GO_AGENT")}, + ArtifactPath: requiredEnv("GO_AGENT"), + }, + } + + encoder := json.NewEncoder(os.Stdout) + encode(encoder, benchmarkMetadata{ + Type: "metadata", + GOOS: runtime.GOOS, + GOARCH: runtime.GOARCH, + Rounds: rounds, + StartupWarmups: startupWarmups, + StartupIterations: startupIterations, + WarmupRequests: warmupRequests, + RPCRequests: rpcRequests, + ManagementCalls: managementRequests, + }) + + for _, result := range benchmarkStartups(agents, startupWarmups, startupIterations) { + encode(encoder, result) + } + + managementURL, closeManagementServer := startManagementServer() + defer closeManagementServer() + managementParams := map[string]any{ + "connection": map[string]any{ + "management_url": managementURL, + "username": "guest", + "password": "guest", + "virtual_host": "/", + }, + } + + for round := 1; round <= rounds; round++ { + order := agents + if round%2 == 0 { + order = []agentSpec{agents[1], agents[0]} + } + for _, agent := range order { + encode(encoder, benchmarkRPC(agent, "handshake", round, "handshake", map[string]any{}, warmupRequests, rpcRequests)) + encode(encoder, benchmarkRPC( + agent, + "management_list_topics", + round, + "mq_list_topics", + managementParams, + warmupRequests/5, + managementRequests, + )) + } + } +} + +func benchmarkStartups(agents []agentSpec, warmups, iterations int) []benchmarkResult { + for warmup := 0; warmup < warmups; warmup++ { + order := agents + if warmup%2 == 1 { + order = []agentSpec{agents[1], agents[0]} + } + for _, agent := range order { + process, _, err := startAgent(agent.Command) + if err != nil { + panic(fmt.Errorf("warm up startup %s: %w", agent.Name, err)) + } + if _, err := process.call("handshake", map[string]any{}); err != nil { + process.kill() + panic(fmt.Errorf("warm up handshake %s: %w", agent.Name, err)) + } + if err := process.close(); err != nil { + panic(fmt.Errorf("close startup warmup %s: %w", agent.Name, err)) + } + } + } + + readySamples := map[string][]float64{} + handshakeSamples := map[string][]float64{} + rssSamples := map[string][]int64{} + readyDurations := map[string]time.Duration{} + handshakeDurations := map[string]time.Duration{} + for iteration := 0; iteration < iterations; iteration++ { + order := agents + if iteration%2 == 1 { + order = []agentSpec{agents[1], agents[0]} + } + for _, agent := range order { + process, readyDuration, err := startAgent(agent.Command) + if err != nil { + panic(fmt.Errorf("start %s: %w", agent.Name, err)) + } + handshakeStart := time.Now() + if _, err := process.call("handshake", map[string]any{}); err != nil { + process.kill() + panic(fmt.Errorf("handshake %s: %w", agent.Name, err)) + } + handshakeDuration := time.Since(handshakeStart) + readySamples[agent.Name] = append(readySamples[agent.Name], milliseconds(readyDuration)) + handshakeSamples[agent.Name] = append( + handshakeSamples[agent.Name], + milliseconds(readyDuration+handshakeDuration), + ) + rssSamples[agent.Name] = append(rssSamples[agent.Name], readRSSKB(process.command.Process.Pid)) + readyDurations[agent.Name] += readyDuration + handshakeDurations[agent.Name] += readyDuration + handshakeDuration + if err := process.close(); err != nil { + panic(fmt.Errorf("close %s: %w", agent.Name, err)) + } + } + } + + results := make([]benchmarkResult, 0, len(agents)*2) + for _, agent := range agents { + artifactBytes := fileSize(agent.ArtifactPath) + ready := summarize(agent.Name, "startup_ready", 0, readySamples[agent.Name], readyDurations[agent.Name], 0) + ready.ReadyRSSKB = medianInt64(rssSamples[agent.Name]) + ready.ArtifactBytes = artifactBytes + results = append(results, ready) + withHandshake := summarize( + agent.Name, + "startup_handshake", + 0, + handshakeSamples[agent.Name], + handshakeDurations[agent.Name], + 0, + ) + withHandshake.ReadyRSSKB = medianInt64(rssSamples[agent.Name]) + withHandshake.ArtifactBytes = artifactBytes + results = append(results, withHandshake) + } + return results +} + +func benchmarkRPC( + agent agentSpec, + workload string, + round int, + method string, + params map[string]any, + warmupRequests int, + operations int, +) benchmarkResult { + process, _, err := startAgent(agent.Command) + if err != nil { + panic(fmt.Errorf("start %s: %w", agent.Name, err)) + } + defer func() { + if err := process.close(); err != nil { + panic(fmt.Errorf("close %s: %w", agent.Name, err)) + } + }() + readyRSS := readRSSKB(process.command.Process.Pid) + for request := 0; request < warmupRequests; request++ { + if _, err := process.call(method, params); err != nil { + panic(fmt.Errorf("warm up %s/%s: %w", agent.Name, workload, err)) + } + } + + latencies := make([]float64, 0, operations) + errorsCount := 0 + start := time.Now() + for operation := 0; operation < operations; operation++ { + requestStart := time.Now() + if _, err := process.call(method, params); err != nil { + errorsCount++ + } + latencies = append(latencies, milliseconds(time.Since(requestStart))) + } + duration := time.Since(start) + result := summarize(agent.Name, workload, round, latencies, duration, errorsCount) + result.ReadyRSSKB = readyRSS + result.PostLoadRSSKB = readRSSKB(process.command.Process.Pid) + result.ArtifactBytes = fileSize(agent.ArtifactPath) + return result +} + +func summarize( + agent string, + workload string, + round int, + latencies []float64, + duration time.Duration, + errorsCount int, +) benchmarkResult { + sorted := append([]float64(nil), latencies...) + sort.Float64s(sorted) + total := 0.0 + for _, latency := range sorted { + total += latency + } + operations := len(sorted) + mean := 0.0 + qps := 0.0 + if operations > 0 { + mean = total / float64(operations) + } + if duration > 0 { + qps = float64(operations) / duration.Seconds() + } + return benchmarkResult{ + Agent: agent, + Workload: workload, + Round: round, + Operations: operations, + Errors: errorsCount, + DurationMS: milliseconds(duration), + QPS: qps, + MeanMS: mean, + P50MS: percentile(sorted, 0.50), + P95MS: percentile(sorted, 0.95), + P99MS: percentile(sorted, 0.99), + } +} + +func startAgent(command []string) (*agentProcess, time.Duration, error) { + if len(command) == 0 { + return nil, 0, errors.New("empty agent command") + } + process := &agentProcess{} + process.command = exec.Command(command[0], command[1:]...) + process.command.Env = sanitizedEnv() + stdin, err := process.command.StdinPipe() + if err != nil { + return nil, 0, err + } + stdout, err := process.command.StdoutPipe() + if err != nil { + return nil, 0, err + } + process.command.Stderr = os.Stderr + process.stdin = stdin + process.reader = bufio.NewScanner(stdout) + process.reader.Buffer(make([]byte, 64*1024), 512*1024*1024) + start := time.Now() + if err := process.command.Start(); err != nil { + return nil, 0, err + } + if !process.reader.Scan() { + process.kill() + return nil, 0, fmt.Errorf("agent did not become ready: %v", process.reader.Err()) + } + if !strings.Contains(process.reader.Text(), `"ready":true`) { + process.kill() + return nil, 0, fmt.Errorf("agent did not become ready: %s", process.reader.Text()) + } + return process, time.Since(start), nil +} + +func (process *agentProcess) call(method string, params map[string]any) (json.RawMessage, error) { + process.nextID++ + request := map[string]any{ + "jsonrpc": "2.0", + "id": process.nextID, + "method": method, + "params": params, + } + payload, err := json.Marshal(request) + if err != nil { + return nil, err + } + if _, err := process.stdin.Write(append(payload, '\n')); err != nil { + return nil, err + } + if !process.reader.Scan() { + return nil, fmt.Errorf("agent response unavailable: %v", process.reader.Err()) + } + var response agentResponse + if err := json.Unmarshal(process.reader.Bytes(), &response); err != nil { + return nil, err + } + if response.ID != process.nextID { + return nil, fmt.Errorf("response id %d does not match request id %d", response.ID, process.nextID) + } + if response.Error != nil { + return nil, errors.New(response.Error.Message) + } + return response.Result, nil +} + +func (process *agentProcess) close() error { + _, callError := process.call("shutdown", map[string]any{}) + _ = process.stdin.Close() + waitError := process.command.Wait() + if callError != nil { + return callError + } + return waitError +} + +func (process *agentProcess) kill() { + if process != nil && process.command != nil && process.command.Process != nil { + _ = process.command.Process.Kill() + _, _ = process.command.Process.Wait() + } +} + +func startManagementServer() (string, func()) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + panic(err) + } + queues := make([]map[string]any, 0, 12) + for index := 11; index >= 0; index-- { + queues = append(queues, map[string]any{ + "name": fmt.Sprintf("queue-%02d", index), + "durable": index%2 == 0, + "auto_delete": index%3 == 0, + "state": "running", + "messages": index * 100, + "consumers": index % 4, + }) + } + body, err := json.Marshal(map[string]any{ + "items": queues, + "page": 1, + "page_count": 1, + "total_count": len(queues), + }) + if err != nil { + panic(err) + } + server := &http.Server{Handler: http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write(body) + })} + go func() { + if err := server.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) { + panic(err) + } + }() + return "http://" + listener.Addr().String(), func() { + _ = server.Close() + } +} + +func javaAgentCommand(jarPath string) []string { + return []string{ + "java", + "-Dfile.encoding=UTF-8", + "-Dsun.stdout.encoding=UTF-8", + "-Dsun.stderr.encoding=UTF-8", + "-Djava.net.useSystemProxies=false", + "-Dhttp.proxyHost=", + "-Dhttps.proxyHost=", + "-DsocksProxyHost=", + "-Doracle.net.disableOob=true", + "-Doracle.jdbc.javaNetNio=false", + "--add-opens=java.sql/java.sql=ALL-UNNAMED", + "-XX:TieredStopAtLevel=1", + "-XX:+UseSerialGC", + "-jar", + jarPath, + } +} + +func sanitizedEnv() []string { + blocked := map[string]struct{}{ + "HTTP_PROXY": {}, "HTTPS_PROXY": {}, "ALL_PROXY": {}, "NO_PROXY": {}, + "http_proxy": {}, "https_proxy": {}, "all_proxy": {}, "no_proxy": {}, + } + result := make([]string, 0, len(os.Environ())) + for _, variable := range os.Environ() { + key, _, _ := strings.Cut(variable, "=") + if _, skip := blocked[key]; !skip { + result = append(result, variable) + } + } + return result +} + +func percentile(sorted []float64, ratio float64) float64 { + if len(sorted) == 0 { + return 0 + } + index := int(math.Ceil(ratio*float64(len(sorted)))) - 1 + if index < 0 { + index = 0 + } + return sorted[index] +} + +func medianInt64(values []int64) int64 { + if len(values) == 0 { + return 0 + } + sorted := append([]int64(nil), values...) + sort.Slice(sorted, func(left, right int) bool { return sorted[left] < sorted[right] }) + return sorted[len(sorted)/2] +} + +func readRSSKB(processID int) int64 { + output, err := exec.Command("ps", "-o", "rss=", "-p", strconv.Itoa(processID)).Output() + if err != nil { + return 0 + } + value, err := strconv.ParseInt(strings.TrimSpace(string(output)), 10, 64) + if err != nil { + return 0 + } + return value +} + +func fileSize(path string) int64 { + info, err := os.Stat(path) + if err != nil { + panic(err) + } + return info.Size() +} + +func milliseconds(duration time.Duration) float64 { + return float64(duration.Nanoseconds()) / float64(time.Millisecond) +} + +func requiredEnv(key string) string { + value := strings.TrimSpace(os.Getenv(key)) + if value == "" { + panic(key + " is required") + } + return value +} + +func envInt(key string, fallback int) int { + value := strings.TrimSpace(os.Getenv(key)) + if value == "" { + return fallback + } + parsed, err := strconv.Atoi(value) + if err != nil || parsed < 1 { + panic(key + " must be a positive integer") + } + return parsed +} + +func encode(encoder *json.Encoder, value any) { + if err := encoder.Encode(value); err != nil { + panic(err) + } +} diff --git a/agents/drivers/rabbitmq/build.gradle b/agents/drivers/rabbitmq/build.gradle deleted file mode 100644 index b1a36133a..000000000 --- a/agents/drivers/rabbitmq/build.gradle +++ /dev/null @@ -1,12 +0,0 @@ -dependencies { - implementation 'com.google.code.gson:gson:2.12.1' - implementation 'com.rabbitmq:amqp-client:5.21.0' - runtimeOnly 'org.slf4j:slf4j-simple:1.7.36' -} - -tasks.named('shadowJar') { - mergeServiceFiles() - manifest { - attributes('Agent-Label': 'RabbitMQ', 'Main-Class': 'com.dbx.agent.rabbitmq.RabbitMqAgent') - } -} diff --git a/agents/drivers/rabbitmq/go.mod b/agents/drivers/rabbitmq/go.mod new file mode 100644 index 000000000..a3769c9a6 --- /dev/null +++ b/agents/drivers/rabbitmq/go.mod @@ -0,0 +1,5 @@ +module github.com/t8y2/dbx/agents/drivers/rabbitmq + +go 1.22 + +require github.com/rabbitmq/amqp091-go v1.13.0 diff --git a/agents/drivers/rabbitmq/go.sum b/agents/drivers/rabbitmq/go.sum new file mode 100644 index 000000000..ecb74926b --- /dev/null +++ b/agents/drivers/rabbitmq/go.sum @@ -0,0 +1,4 @@ +github.com/rabbitmq/amqp091-go v1.13.0 h1:L8NA1WtF76C6KA3LAoufjfLgbist/If1UQYcsOjtxXA= +github.com/rabbitmq/amqp091-go v1.13.0/go.mod h1:Hy4jKW5kQART1u+JkDTF9YYOQUHXqMuhrgxOEeS7G4o= +go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= +go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= diff --git a/agents/drivers/rabbitmq/helpers.go b/agents/drivers/rabbitmq/helpers.go new file mode 100644 index 000000000..fb2f8e26a --- /dev/null +++ b/agents/drivers/rabbitmq/helpers.go @@ -0,0 +1,448 @@ +package main + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "regexp" + "strconv" + "strings" + "time" + + amqp "github.com/rabbitmq/amqp091-go" +) + +var ( + quotedNamePattern = regexp.MustCompile(`'([^']+)'`) + declaredResourcePattern = regexp.MustCompile(`for (queue|exchange) '([^']+)'`) +) + +func decodeJSON(data []byte, target any) error { + decoder := json.NewDecoder(bytes.NewReader(data)) + decoder.UseNumber() + return decoder.Decode(target) +} + +func deepCopyObject(source jsonObject) jsonObject { + if source == nil { + return nil + } + encoded, err := json.Marshal(source) + if err != nil { + return nil + } + copy := jsonObject{} + if err := decodeJSON(encoded, ©); err != nil { + return nil + } + return copy +} + +func okResult() jsonObject { + return jsonObject{"ok": true} +} + +func objectOrNil(object jsonObject, key string) jsonObject { + if object == nil { + return nil + } + switch value := object[key].(type) { + case jsonObject: + return value + case map[string]any: + return jsonObject(value) + default: + return nil + } +} + +func arrayOrNil(object jsonObject, key string) []any { + if object == nil { + return nil + } + array, _ := object[key].([]any) + return array +} + +func stringOrNull(object jsonObject, key string) *string { + if object == nil { + return nil + } + value, exists := object[key] + if !exists || value == nil { + return nil + } + var result string + switch typed := value.(type) { + case string: + result = typed + case json.Number: + result = typed.String() + case bool: + result = strconv.FormatBool(typed) + case float64: + result = strconv.FormatFloat(typed, 'f', -1, 64) + default: + result = fmt.Sprint(typed) + } + return &result +} + +func stringOrEmpty(object jsonObject, key string) string { + return stringOrDefault(object, key, "") +} + +func stringOrDefault(object jsonObject, key, fallback string) string { + value := stringOrNull(object, key) + if value == nil { + return fallback + } + return *value +} + +func integerOrNull(object jsonObject, key string) *int { + value, ok := numberAsInt64(object, key) + if !ok { + return nil + } + converted := int(value) + return &converted +} + +func longOrNull(object jsonObject, key string) *int64 { + value, ok := numberAsInt64(object, key) + if !ok { + return nil + } + return &value +} + +func numberAsInt64(object jsonObject, key string) (int64, bool) { + if object == nil { + return 0, false + } + value, exists := object[key] + if !exists || value == nil { + return 0, false + } + switch typed := value.(type) { + case json.Number: + if integer, err := typed.Int64(); err == nil { + return integer, true + } + decimal, err := typed.Float64() + return int64(decimal), err == nil + case float64: + return int64(typed), true + case float32: + return int64(typed), true + case int: + return int64(typed), true + case int8: + return int64(typed), true + case int16: + return int64(typed), true + case int32: + return int64(typed), true + case int64: + return typed, true + case uint: + return int64(typed), true + case uint8: + return int64(typed), true + case uint16: + return int64(typed), true + case uint32: + return int64(typed), true + case uint64: + return int64(typed), true + case string: + integer, err := strconv.ParseInt(typed, 10, 64) + return integer, err == nil + default: + return 0, false + } +} + +func intOrDefault(object jsonObject, key string, fallback int) int { + value := integerOrNull(object, key) + if value == nil { + return fallback + } + return *value +} + +func longOrDefault(object jsonObject, key string, fallback int64) int64 { + value := longOrNull(object, key) + if value == nil { + return fallback + } + return *value +} + +func floatOrNull(object jsonObject, key string) *float64 { + if object == nil { + return nil + } + value, exists := object[key] + if !exists || value == nil { + return nil + } + var result float64 + var err error + switch typed := value.(type) { + case json.Number: + result, err = typed.Float64() + case float64: + result = typed + case float32: + result = float64(typed) + case int: + result = float64(typed) + case int64: + result = float64(typed) + case string: + result, err = strconv.ParseFloat(typed, 64) + default: + return nil + } + if err != nil { + return nil + } + return &result +} + +func boolOrDefault(object jsonObject, key string, fallback bool) bool { + if object == nil { + return fallback + } + value, exists := object[key] + if !exists || value == nil { + return fallback + } + switch typed := value.(type) { + case bool: + return typed + case string: + parsed, err := strconv.ParseBool(typed) + if err == nil { + return parsed + } + } + return fallback +} + +func integerProperty(properties jsonObject, key string) (int, bool) { + if properties == nil { + return 0, false + } + value := integerOrNull(properties, key) + if value == nil { + return 0, false + } + return *value, true +} + +func boolProperty(config jsonObject, key string) bool { + return boolOrDefault(objectOrNil(config, "properties"), key, false) +} + +func durationMilliseconds(object jsonObject, key string, fallback time.Duration) time.Duration { + value := integerOrNull(object, key) + if value == nil { + return fallback + } + return time.Duration(*value) * time.Millisecond +} + +func argumentValue(value any) any { + switch typed := value.(type) { + case nil: + return nil + case bool, string: + return typed + case json.Number: + if integer, err := typed.Int64(); err == nil { + return integer + } + if decimal, err := typed.Float64(); err == nil { + return int64(decimal) + } + return nil + case float64: + return int64(typed) + case float32: + return int64(typed) + case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64: + return typed + default: + return nil + } +} + +func connectionObject(params jsonObject) jsonObject { + if connection := objectOrNil(params, "connection"); connection != nil { + return connection + } + return params +} + +func (s *server) currentConnectionConfig(params jsonObject) jsonObject { + if connection := objectOrNil(params, "connection"); connection != nil { + return connection + } + return s.cachedConnection +} + +func (s *server) requireConnectionConfig(params jsonObject) (jsonObject, error) { + connection := s.currentConnectionConfig(params) + if connection == nil { + return nil, errors.New("Not connected. Call connect first.") + } + return connection, nil +} + +func (s *server) requireConnection() (*amqp.Connection, error) { + if s.connection == nil { + return nil, errors.New("Not connected. Call connect first.") + } + return s.connection, nil +} + +func queueName(params jsonObject) (string, error) { + name := stringOrEmpty(params, "topic") + if strings.TrimSpace(name) == "" { + name = stringOrEmpty(params, "name") + } + if strings.TrimSpace(name) == "" { + return "", errors.New("topic (queue name) is required") + } + return name, nil +} + +func effectiveVhost(params, connection jsonObject) string { + vhost := stringOrNull(params, "virtual_host") + if vhost == nil || strings.TrimSpace(*vhost) == "" { + if connection != nil { + return stringOrDefault(connection, "virtual_host", "/") + } + return "/" + } + return *vhost +} + +func allVhostsRequested(params jsonObject) bool { + return boolOrDefault(params, "all_vhosts", false) +} + +func managementListPath(params, connection jsonObject, resource string) string { + if allVhostsRequested(params) { + return "/api/" + resource + } + return "/api/" + resource + "/" + urlEncodeVhost(effectiveVhost(params, connection)) +} + +func vhostFilter(params, connection jsonObject) string { + if allVhostsRequested(params) { + return "" + } + return effectiveVhost(params, connection) +} + +func attachVhost(info jsonObject, source jsonObject) { + info["vhost"] = stringOrEmpty(source, "vhost") +} + +func serverString(properties amqp.Table, key string) any { + value, exists := properties[key] + if !exists || value == nil { + return nil + } + switch typed := value.(type) { + case []byte: + return string(typed) + default: + return fmt.Sprint(typed) + } +} + +func normalizeErrorMessage(err error) string { + if err == nil { + return "error" + } + if errors.Is(err, amqp.ErrCredentials) { + return err.Error() + ". Hint: authentication failed. Check the RabbitMQ username, password, and virtual host permissions." + } + var amqpError *amqp.Error + if errors.As(err, &amqpError) { + if friendly := mapAMQPError(amqpError.Code, amqpError.Reason); friendly != "" { + return friendly + } + } + message := strings.TrimSpace(err.Error()) + if message == "" { + return fmt.Sprintf("%T", err) + } + return message +} + +func mapAMQPError(replyCode int, replyText string) string { + switch replyCode { + case 405: + name := extractQuotedName(replyText) + subject := "The queue" + if name != "" { + subject = "Queue '" + name + "'" + } + return subject + " is exclusive and owned by another connection. Hint: exclusive queues can only be accessed by their owning connection; stats via the management API are still available." + case 404: + name := extractQuotedName(replyText) + kind := "Queue" + if strings.Contains(replyText, "no exchange") { + kind = "Exchange" + } + subject := "The " + strings.ToLower(kind) + if name != "" { + subject = kind + " '" + name + "'" + } + return subject + " was not found. Hint: it may have been deleted, or it never existed on this virtual host." + case 406: + name := extractDeclaredResourceName(replyText) + kind := "Queue" + if strings.Contains(replyText, "for exchange") { + kind = "Exchange" + } + subject := "The " + strings.ToLower(kind) + if name != "" { + subject = kind + " '" + name + "'" + } + lowerKind := strings.ToLower(kind) + return subject + " already exists with different parameters. Hint: " + lowerKind + " parameters are immutable after declaration; delete and re-declare the " + lowerKind + " to change them." + case 403: + name := extractQuotedName(replyText) + subject := "the requested resource" + if name != "" { + subject = "'" + name + "'" + } + return "Access to " + subject + " was refused. Hint: check the user's configure/write/read permissions on the virtual host." + default: + return "" + } +} + +func extractQuotedName(replyText string) string { + match := quotedNamePattern.FindStringSubmatch(replyText) + if len(match) < 2 { + return "" + } + return match[1] +} + +func extractDeclaredResourceName(replyText string) string { + match := declaredResourcePattern.FindStringSubmatch(replyText) + if len(match) < 3 { + return "" + } + return match[2] +} diff --git a/agents/drivers/rabbitmq/helpers_test.go b/agents/drivers/rabbitmq/helpers_test.go new file mode 100644 index 000000000..98b4d2f2e --- /dev/null +++ b/agents/drivers/rabbitmq/helpers_test.go @@ -0,0 +1,305 @@ +package main + +import ( + "encoding/base64" + "encoding/json" + "strings" + "testing" +) + +func mustObject(t *testing.T, source string) jsonObject { + t.Helper() + result := jsonObject{} + if err := decodeJSON([]byte(source), &result); err != nil { + t.Fatal(err) + } + return result +} + +func TestParseAddresses(t *testing.T) { + tests := []struct { + name string + value string + defaultPort int + want []address + wantError string + }{ + {name: "pairs", value: "host1:5672,host2:5673", defaultPort: 5672, want: []address{{"host1", 5672}, {"host2", 5673}}}, + {name: "bare host", value: "rabbit", defaultPort: 5679, want: []address{{"rabbit", 5679}}}, + {name: "blank entries", value: " , rabbit:5672, ", defaultPort: 5679, want: []address{{"rabbit", 5672}}}, + {name: "ipv6", value: "[::1]:5672", defaultPort: 5679, want: []address{{"::1", 5672}}}, + {name: "blank", value: " , ", defaultPort: 5672, wantError: "addresses is required"}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got, err := parseAddresses(test.value, test.defaultPort) + if test.wantError != "" { + if err == nil || !strings.Contains(err.Error(), test.wantError) { + t.Fatalf("got error %v, want %q", err, test.wantError) + } + return + } + if err != nil { + t.Fatal(err) + } + if len(got) != len(test.want) { + t.Fatalf("got %#v, want %#v", got, test.want) + } + for index := range got { + if got[index] != test.want[index] { + t.Fatalf("got %#v, want %#v", got, test.want) + } + } + }) + } +} + +func TestResolveAddresses(t *testing.T) { + tests := []struct { + name string + config jsonObject + want []address + wantError string + }{ + {name: "explicit port", config: mustObject(t, `{"addresses":"rabbit","port":5679}`), want: []address{{"rabbit", 5679}}}, + {name: "default port", config: mustObject(t, `{"addresses":"rabbit"}`), want: []address{{"rabbit", 5672}}}, + {name: "host fallback", config: mustObject(t, `{"host":"rabbit"}`), want: []address{{"rabbit", 5672}}}, + {name: "missing", config: mustObject(t, `{}`), wantError: "addresses is required"}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got, err := resolveAddresses(test.config) + if test.wantError != "" { + if err == nil || !strings.Contains(err.Error(), test.wantError) { + t.Fatalf("got error %v, want %q", err, test.wantError) + } + return + } + if err != nil { + t.Fatal(err) + } + if len(got) != len(test.want) || got[0] != test.want[0] { + t.Fatalf("got %#v, want %#v", got, test.want) + } + }) + } +} + +func TestPeekNormalizationAndRoutingKey(t *testing.T) { + if normalizePeekOffset(-1) != 0 || normalizePeekOffset(4) != 4 { + t.Fatal("unexpected offset normalization") + } + if normalizePeekCount(0) != 1 || normalizePeekCount(8) != 8 { + t.Fatal("unexpected count normalization") + } + tests := []struct { + params jsonObject + want string + }{ + {mustObject(t, `{"routing_key":"explicit","routingKey":"camel","key":"message"}`), "explicit"}, + {mustObject(t, `{"routingKey":"camel","key":"message"}`), "camel"}, + {mustObject(t, `{"key":"message"}`), "message"}, + {mustObject(t, `{"key":" "}`), "queue"}, + {mustObject(t, `{}`), "queue"}, + } + for _, test := range tests { + if got := resolveRoutingKey(test.params, "queue"); got != test.want { + t.Fatalf("got %q, want %q", got, test.want) + } + } + if got := peekMessageCapacity(10, 3, int(^uint(0)>>1)); got != 7 { + t.Fatalf("unexpected bounded capacity %d", got) + } + if got := peekMessageCapacity(2, 5, 10); got != 0 { + t.Fatalf("unexpected exhausted capacity %d", got) + } +} + +func TestTLSAndManagementConfiguration(t *testing.T) { + if tlsSkipVerify(mustObject(t, `{"tls_skip_verify":true}`)) != true { + t.Fatal("top-level skip verify not detected") + } + if tlsSkipVerify(mustObject(t, `{"tls":{"skip_verify":true}}`)) != true { + t.Fatal("nested skip verify not detected") + } + if managementTLS(mustObject(t, `{"tls_skip_verify":true}`)) { + t.Fatal("skip verify must not enable management TLS") + } + if !managementTLS(mustObject(t, `{"tls":{}}`)) || !managementTLS(mustObject(t, `{"properties":{"ssl":true}}`)) { + t.Fatal("management TLS not detected") + } + if managementPort(jsonObject{}, false) != 15672 || managementPort(jsonObject{}, true) != 15671 { + t.Fatal("unexpected default management ports") + } + if managementPort(mustObject(t, `{"properties":{"management_port":55672}}`), false) != 55672 { + t.Fatal("management port override ignored") + } +} + +func TestCredentialAndAuthHelpers(t *testing.T) { + config := mustObject(t, `{"username":" ","password":null}`) + if credentialOrGuest(config, "username") != "guest" || credentialOrGuest(config, "password") != "guest" { + t.Fatal("blank credentials did not fall back to guest") + } + want := "Basic " + base64.StdEncoding.EncodeToString([]byte("guest:guest")) + if got := basicAuthHeader("guest", "guest"); got != want { + t.Fatalf("got %q, want %q", got, want) + } +} + +func TestPathEncoding(t *testing.T) { + if got := urlEncodeVhost("/"); got != "%2F" { + t.Fatalf("got %q", got) + } + if got := urlEncodePathSegment("queue one"); got != "queue%20one" { + t.Fatalf("got %q", got) + } + if got := urlEncodeName("127.0.0.1:1 -> 127.0.0.1:2"); !strings.Contains(got, "%20-%3E%20") { + t.Fatalf("got %q", got) + } +} + +func TestHandshakeAndRequestErrors(t *testing.T) { + service := newServer() + response, shutdown := service.handleRequest([]byte(`{"jsonrpc":"2.0","id":1,"method":"handshake","params":{}}`)) + if shutdown || response.Error != nil { + t.Fatalf("unexpected response: %#v", response) + } + result, ok := response.Result.(handshakeResult) + if !ok || result.ProtocolVersion != 1 || result.AgentProtocolVersion != 1 || len(result.Capabilities) != len(capabilities) { + t.Fatalf("unexpected handshake: %#v", response.Result) + } + response, _ = service.handleRequest([]byte(`{"jsonrpc":"2.0","id":2,"method":"unknown","params":{}}`)) + if response.Error == nil || !strings.Contains(response.Error.Message, "Unknown method") { + t.Fatalf("unexpected response: %#v", response) + } + response, _ = service.handleRequest([]byte(`not json`)) + if response.Error == nil || string(response.ID) != "null" { + t.Fatalf("unexpected malformed response: %#v", response) + } + response, _ = service.handleRequest([]byte(`{"jsonrpc":"2.0","id":7,"params":{}}`)) + if response.Error == nil || string(response.ID) != "7" { + t.Fatalf("unexpected missing-method response: %#v", response) + } + encoded, err := json.Marshal(response) + if err != nil || !strings.Contains(string(encoded), `"id":7`) { + t.Fatalf("unexpected JSON: %s, %v", encoded, err) + } +} + +func TestAllVhostsGuardsAndEffectiveVhost(t *testing.T) { + service := newServer() + for method := range allVhostsUnsupportedMethods { + _, _, err := service.dispatch(method, mustObject(t, `{"all_vhosts":true}`)) + if err == nil || err.Error() != "all_vhosts is only supported for list operations" { + t.Fatalf("%s: %v", method, err) + } + } + connection := mustObject(t, `{"virtual_host":"connected"}`) + if got := effectiveVhost(mustObject(t, `{"virtual_host":"explicit"}`), connection); got != "explicit" { + t.Fatalf("got %q", got) + } + if got := effectiveVhost(mustObject(t, `{"virtual_host":" "}`), connection); got != "connected" { + t.Fatalf("got %q", got) + } + if got := effectiveVhost(jsonObject{}, nil); got != "/" { + t.Fatalf("got %q", got) + } + if allVhostsRequested(jsonObject{}) { + t.Fatal("all_vhosts should default false") + } + if got := managementListPath(jsonObject{}, connection, "queues"); got != "/api/queues/connected" { + t.Fatalf("got %q", got) + } + if got := managementListPath(mustObject(t, `{"all_vhosts":true,"virtual_host":"ignored"}`), connection, "queues"); got != "/api/queues" { + t.Fatalf("got %q", got) + } + if got := vhostFilter(mustObject(t, `{"all_vhosts":true}`), connection); got != "" { + t.Fatalf("got %q", got) + } +} + +func TestSemanticGuards(t *testing.T) { + if _, err := queueName(jsonObject{}); err == nil || !strings.Contains(err.Error(), "queue name") { + t.Fatal(err) + } + if _, err := namespaceName(jsonObject{}); err == nil || err.Error() != "namespace is required" { + t.Fatal(err) + } + if _, err := namespaceName(mustObject(t, `{"namespace":"*"}`)); err == nil { + t.Fatal("all-vhosts namespace accepted") + } + if err := assertNamespaceDeletable("/", ""); err == nil { + t.Fatal("default vhost deletion accepted") + } + if err := assertNamespaceDeletable("orders", "orders"); err == nil { + t.Fatal("connected vhost deletion accepted") + } + for _, exchangeType := range []string{"direct", "fanout", "topic", "headers"} { + if _, err := validateExchangeType(exchangeType); err != nil { + t.Fatal(err) + } + } + if _, err := validateExchangeType("stream"); err == nil { + t.Fatal("invalid exchange type accepted") + } + for _, name := range []string{"", "amq.direct"} { + if err := assertExchangeDeletable(name); err == nil { + t.Fatalf("exchange %q accepted", name) + } + } + if _, err := permissionVhost(jsonObject{}); err == nil { + t.Fatal("blank permission vhost accepted") + } + if _, err := permissionVhost(mustObject(t, `{"virtual_host":"*"}`)); err == nil { + t.Fatal("all-vhosts permission accepted") + } + if err := assertNotConnectedUser("delete", "dbx", "dbx"); err == nil { + t.Fatal("connected user mutation accepted") + } +} + +func TestPermissionAndUserHelpers(t *testing.T) { + if permissionPattern(jsonObject{}, "read") != ".*" || permissionPattern(mustObject(t, `{"read":"^q"}`), "read") != "^q" { + t.Fatal("unexpected permission pattern") + } + if got := parseUserTags("administrator, management, ,policymaker"); len(got) != 3 || got[1] != "management" { + t.Fatalf("got %#v", got) + } + if got := userTagsParam(mustObject(t, `{"tags":["management"," policymaker ",""]}`)); got != "management,policymaker" { + t.Fatalf("got %q", got) + } + if got := userTagsParam(mustObject(t, `{"tags":"administrator,management"}`)); got != "administrator,management" { + t.Fatalf("got %q", got) + } +} + +func TestAMQPErrorMapping(t *testing.T) { + tests := []struct { + code int + text string + want string + }{ + {405, "RESOURCE_LOCKED - cannot obtain exclusive access to locked queue 'q1'", "Queue 'q1' is exclusive"}, + {405, "RESOURCE_LOCKED", "The queue is exclusive"}, + {404, "NOT_FOUND - no queue 'q1' in vhost '/'", "Queue 'q1' was not found"}, + {404, "NOT_FOUND - no exchange 'events' in vhost '/'", "Exchange 'events' was not found"}, + {406, "PRECONDITION_FAILED - inequivalent arg 'durable' for queue 'q1' in vhost '/'", "Queue 'q1' already exists"}, + {406, "PRECONDITION_FAILED - inequivalent arg 'type' for exchange 'events' in vhost '/'", "Exchange 'events' already exists"}, + {403, "ACCESS_REFUSED - access to queue 'q1' refused", "Access to 'q1' was refused"}, + } + for _, test := range tests { + if got := mapAMQPError(test.code, test.text); !strings.Contains(got, test.want) { + t.Fatalf("got %q, want substring %q", got, test.want) + } + } + if got := mapAMQPError(320, "CONNECTION_FORCED"); got != "" { + t.Fatalf("unexpected mapping %q", got) + } + if got := extractDeclaredResourceName("inequivalent arg 'durable' for queue 'q1' in vhost '/'"); got != "q1" { + t.Fatalf("got %q", got) + } + if got := extractQuotedName("access to queue 'q1' refused for user 'dbx'"); got != "q1" { + t.Fatalf("got %q", got) + } +} diff --git a/agents/drivers/rabbitmq/integration_test.go b/agents/drivers/rabbitmq/integration_test.go new file mode 100644 index 000000000..2ad2c8d21 --- /dev/null +++ b/agents/drivers/rabbitmq/integration_test.go @@ -0,0 +1,248 @@ +package main + +import ( + "encoding/base64" + "fmt" + "os" + "strconv" + "strings" + "testing" + "time" +) + +func TestRabbitMQIntegration(t *testing.T) { + if os.Getenv("RABBITMQ_INTEGRATION") != "1" { + t.Skip("set RABBITMQ_INTEGRATION=1 to run against a real RabbitMQ broker") + } + host := envOrDefault("RABBITMQ_HOST", "127.0.0.1") + amqpPort := envIntOrDefault(t, "RABBITMQ_PORT", 5672) + managementPort := envIntOrDefault(t, "RABBITMQ_MANAGEMENT_PORT", 15672) + username := envOrDefault("RABBITMQ_USERNAME", "dbx") + password := envOrDefault("RABBITMQ_PASSWORD", "dbx-password") + connection := jsonObject{ + "addresses": host, + "port": amqpPort, + "username": username, + "password": password, + "properties": jsonObject{ + "management_port": managementPort, + }, + } + service := newServer() + if _, err := service.connect(connection); err != nil { + t.Fatal(err) + } + defer service.closeClients() + + probe, err := service.testConnection(connection) + if err != nil { + t.Fatal(err) + } + if probe.(jsonObject)["ok"] != true || probe.(jsonObject)["serverVersion"] == nil { + t.Fatalf("unexpected probe %#v", probe) + } + badConnection := deepCopyObject(connection) + badConnection["password"] = "definitely-wrong-password" + if _, err := service.testConnection(badConnection); err == nil { + t.Fatal("expected invalid credentials to fail") + } else if !strings.Contains(normalizeErrorMessage(err), "authentication failed") { + t.Fatalf("authentication error lost its actionable hint: %v", err) + } + + suffix := fmt.Sprintf("%d", time.Now().UnixNano()) + vhost := "dbx-go-" + suffix + queue := "queue-" + suffix + exchange := "exchange-" + suffix + policy := "policy-" + suffix + user := "user-" + suffix + + defer managementSend(connection, "DELETE", "/api/users/"+urlEncodePathSegment(user), nil) + defer managementSend(connection, "DELETE", "/api/vhosts/"+urlEncodeVhost(vhost), nil) + + if _, err := service.createNamespace(jsonObject{"namespace": vhost}); err != nil { + t.Fatal(err) + } + if _, err := service.grantPermission(jsonObject{ + "user": username, "virtual_host": vhost, "configure": ".*", "write": ".*", "read": ".*", + }); err != nil { + t.Fatal(err) + } + if _, err := service.getTopicStats(jsonObject{"topic": "missing-" + suffix, "virtual_host": vhost}); err == nil { + t.Fatal("expected missing queue lookup to fail") + } else if !strings.Contains(normalizeErrorMessage(err), "was not found") { + t.Fatalf("unexpected missing queue error: %v", err) + } + if _, err := service.createTopic(jsonObject{"topic": queue, "virtual_host": vhost, "durable": true}); err != nil { + t.Fatalf("channel did not recover after a broker-forced close: %v", err) + } + if _, err := service.createExchange(jsonObject{ + "name": exchange, "type": "topic", "virtual_host": vhost, "durable": true, + }); err != nil { + t.Fatal(err) + } + if _, err := service.bind(jsonObject{ + "source": exchange, "destination": queue, "destinationType": "queue", "routingKey": "orders.*", "virtual_host": vhost, + }); err != nil { + t.Fatal(err) + } + payload := "RabbitMQ Go Agent 世界" + if _, err := service.sendMessage(jsonObject{ + "topic": queue, "exchange": exchange, "routingKey": "orders.created", "virtual_host": vhost, + "payloadBase64": base64.StdEncoding.EncodeToString([]byte(payload)), + "headers": jsonObject{"source": "integration", "attempt": 1}, + }); err != nil { + t.Fatal(err) + } + + peeked, err := service.peekMessages(jsonObject{"topic": queue, "virtual_host": vhost, "offset": 0, "count": 10}) + if err != nil { + t.Fatal(err) + } + messages := peeked.(jsonObject)["messages"].([]jsonObject) + if len(messages) != 1 || messages[0]["payloadText"] != payload || messages[0]["routingKey"] != "orders.created" { + t.Fatalf("unexpected messages %#v", messages) + } + + stats, err := service.getTopicStats(jsonObject{"topic": queue, "virtual_host": vhost}) + if err != nil { + t.Fatal(err) + } + if stats.(jsonObject)["totalMessages"] != int64(1) { + t.Fatalf("unexpected stats %#v", stats) + } + config, err := service.getTopicConfig(jsonObject{"topic": queue, "virtual_host": vhost}) + if err != nil { + t.Fatal(err) + } + if config.(jsonObject)["configs"].(jsonObject)["durable"] != true { + t.Fatalf("unexpected config %#v", config) + } + consumers, err := service.listConsumers(jsonObject{"topic": queue, "virtual_host": vhost}) + if err != nil || len(consumers.(jsonObject)["consumers"].([]jsonObject)) != 0 { + t.Fatalf("unexpected consumers %#v, %v", consumers, err) + } + + topics, err := service.listTopics(jsonObject{"virtual_host": vhost}) + if err != nil || !containsNamedItem(topics.(jsonObject)["topics"].([]jsonObject), queue) { + t.Fatalf("unexpected topics %#v, %v", topics, err) + } + exchanges, err := service.listExchanges(jsonObject{"virtual_host": vhost}) + if err != nil || !containsNamedItem(exchanges.(jsonObject)["exchanges"].([]jsonObject), exchange) { + t.Fatalf("unexpected exchanges %#v, %v", exchanges, err) + } + bindings, err := service.listBindings(jsonObject{"virtual_host": vhost, "queue": queue}) + if err != nil || len(bindings.(jsonObject)["bindings"].([]jsonObject)) == 0 { + t.Fatalf("unexpected bindings %#v, %v", bindings, err) + } + + if _, err := service.setPolicy(jsonObject{ + "virtual_host": vhost, "name": policy, "pattern": "^" + queue + "$", "applyTo": "queues", + "definition": jsonObject{"max-length": 1000}, + }); err != nil { + t.Fatal(err) + } + policies, err := service.listPolicies(jsonObject{"virtual_host": vhost}) + if err != nil || !containsNamedItem(policies.(jsonObject)["policies"].([]jsonObject), policy) { + t.Fatalf("unexpected policies %#v, %v", policies, err) + } + if _, err := service.deletePolicy(jsonObject{"virtual_host": vhost, "name": policy}); err != nil { + t.Fatal(err) + } + + if _, err := service.createUser(jsonObject{"name": user, "password": "temporary-password", "tags": []any{"management"}}); err != nil { + t.Fatal(err) + } + users, err := service.listUsers(jsonObject{}) + if err != nil || !containsNamedItem(users.(jsonObject)["users"].([]jsonObject), user) { + t.Fatalf("unexpected users %#v, %v", users, err) + } + if _, err := service.grantPermission(jsonObject{"user": user, "virtual_host": vhost}); err != nil { + t.Fatal(err) + } + permissions, err := service.listPermissions(jsonObject{"user": user, "virtual_host": vhost}) + if err != nil || len(permissions.(jsonObject)["permissions"].([]jsonObject)) != 1 { + t.Fatalf("unexpected permissions %#v, %v", permissions, err) + } + if _, err := service.revokePermission(jsonObject{"user": user, "virtual_host": vhost}); err != nil { + t.Fatal(err) + } + if _, err := service.deleteUser(jsonObject{"name": user}); err != nil { + t.Fatal(err) + } + + namespaces, err := service.listNamespaces(jsonObject{}) + if err != nil || !containsNamedItem(namespaces.(jsonObject)["namespaces"].([]jsonObject), vhost) { + t.Fatalf("unexpected namespaces %#v, %v", namespaces, err) + } + connections, err := service.listClientConnections(jsonObject{"all_vhosts": true}) + if err != nil || len(connections.(jsonObject)["connections"].([]jsonObject)) == 0 { + t.Fatalf("unexpected connections %#v, %v", connections, err) + } + channels, err := service.listClientChannels(jsonObject{"all_vhosts": true}) + if err != nil || len(channels.(jsonObject)["channels"].([]jsonObject)) == 0 { + t.Fatalf("unexpected channels %#v, %v", channels, err) + } + if _, err := service.getOverview(jsonObject{}); err != nil { + t.Fatal(err) + } + nodes, err := service.listNodes(jsonObject{}) + if err != nil || len(nodes.(jsonObject)["nodes"].([]jsonObject)) == 0 { + t.Fatalf("unexpected nodes %#v, %v", nodes, err) + } + cluster, err := service.describeCluster(jsonObject{}) + if err != nil || cluster.(jsonObject)["version"] == nil { + t.Fatalf("unexpected cluster %#v, %v", cluster, err) + } + + purged, err := service.purgeQueue(jsonObject{"topic": queue, "virtual_host": vhost}) + if err != nil || purged.(jsonObject)["purged"] != 1 { + t.Fatalf("unexpected purge %#v, %v", purged, err) + } + if _, err := service.unbind(jsonObject{ + "source": exchange, "destination": queue, "destinationType": "queue", "routingKey": "orders.*", "virtual_host": vhost, + }); err != nil { + t.Fatal(err) + } + if _, err := service.deleteTopic(jsonObject{"topic": queue, "virtual_host": vhost}); err != nil { + t.Fatal(err) + } + if _, err := service.deleteExchange(jsonObject{"name": exchange, "virtual_host": vhost}); err != nil { + t.Fatal(err) + } + if client := service.vhostClients[vhost]; client != nil { + client.close() + delete(service.vhostClients, vhost) + } + if _, err := service.deleteNamespace(jsonObject{"namespace": vhost}); err != nil { + t.Fatal(err) + } +} + +func containsNamedItem(items []jsonObject, name string) bool { + for _, item := range items { + if stringOrEmpty(item, "name") == name { + return true + } + } + return false +} + +func envOrDefault(key, fallback string) string { + if value := strings.TrimSpace(os.Getenv(key)); value != "" { + return value + } + return fallback +} + +func envIntOrDefault(t *testing.T, key string, fallback int) int { + t.Helper() + value := strings.TrimSpace(os.Getenv(key)) + if value == "" { + return fallback + } + parsed, err := strconv.Atoi(value) + if err != nil { + t.Fatalf("invalid %s: %v", key, err) + } + return parsed +} diff --git a/agents/drivers/rabbitmq/main.go b/agents/drivers/rabbitmq/main.go new file mode 100644 index 000000000..bab7c4a3f --- /dev/null +++ b/agents/drivers/rabbitmq/main.go @@ -0,0 +1,529 @@ +package main + +import ( + "bufio" + "crypto/tls" + "encoding/json" + "errors" + "fmt" + "net" + "net/url" + "os" + "strconv" + "strings" + "time" + + amqp "github.com/rabbitmq/amqp091-go" +) + +const ( + protocolVersion = 1 + agentProtocolVersion = 1 + defaultAMQPPort = 5672 + defaultRequestTimeout = 30 * time.Second + defaultHandshakeTimeout = 10 * time.Second + defaultHeartbeat = 60 * time.Second + defaultChannelMax = 2047 + maxRPCMessageBytes = 32 * 1024 * 1024 +) + +var capabilities = []string{ + "mq_connect", "mq_test_connection", "mq_topics", + "mq_messages", "mq_config", "mq_monitoring", "mq_exchanges", + "mq_client_connections", "mq_user_permissions", "mq_policies", +} + +var allVhostsUnsupportedMethods = map[string]struct{}{ + "mq_create_topic": {}, "mq_delete_topic": {}, "mq_purge_queue": {}, "mq_send_message": {}, + "mq_bind": {}, "mq_unbind": {}, "mq_create_exchange": {}, "mq_delete_exchange": {}, + "mq_peek_messages": {}, "mq_get_topic_stats": {}, "mq_list_consumers": {}, "mq_close_connection": {}, + "mq_grant_permission": {}, "mq_revoke_permission": {}, "mq_set_policy": {}, "mq_delete_policy": {}, +} + +type jsonObject map[string]any + +type rpcRequest struct { + ID json.RawMessage `json:"id"` + Method string `json:"method"` + Params json.RawMessage `json:"params"` +} + +type rpcError struct { + Code int `json:"code"` + Message string `json:"message"` +} + +type rpcResponse struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Result any `json:"result,omitempty"` + Error *rpcError `json:"error,omitempty"` +} + +type handshakeResult struct { + ProtocolVersion int `json:"protocolVersion"` + AgentProtocolVersion int `json:"agentProtocolVersion"` + Capabilities []string `json:"capabilities"` +} + +type address struct { + Host string + Port int +} + +type vhostClient struct { + connection *amqp.Connection + channel *amqp.Channel +} + +type server struct { + connection *amqp.Connection + channel *amqp.Channel + cachedConnection jsonObject + vhostClients map[string]*vhostClient +} + +func main() { + service := newServer() + encoder := json.NewEncoder(os.Stdout) + encoder.SetEscapeHTML(false) + fmt.Fprintln(os.Stdout, `{"ready":true}`) + + scanner := bufio.NewScanner(os.Stdin) + scanner.Buffer(make([]byte, 64*1024), maxRPCMessageBytes) + for scanner.Scan() { + response, shutdown := service.handleRequest(scanner.Bytes()) + if err := encoder.Encode(response); err != nil { + fmt.Fprintln(os.Stderr, err) + return + } + if shutdown { + return + } + } + service.closeClients() +} + +func newServer() *server { + return &server{vhostClients: make(map[string]*vhostClient)} +} + +func (s *server) handleRequest(line []byte) (rpcResponse, bool) { + response := rpcResponse{JSONRPC: "2.0", ID: json.RawMessage("null")} + var request rpcRequest + if err := json.Unmarshal(line, &request); err != nil { + response.Error = &rpcError{Code: -1, Message: normalizeErrorMessage(err)} + return response, false + } + if len(request.ID) > 0 { + response.ID = request.ID + } + params := jsonObject{} + if len(request.Params) > 0 && string(request.Params) != "null" { + var decoded any + if err := decodeJSON(request.Params, &decoded); err != nil { + response.Error = &rpcError{Code: -1, Message: normalizeErrorMessage(err)} + return response, false + } + if object, ok := decoded.(map[string]any); ok { + params = jsonObject(object) + } + } + result, shutdown, err := s.dispatch(request.Method, params) + if err != nil { + response.Error = &rpcError{Code: -1, Message: normalizeErrorMessage(err)} + return response, false + } + response.Result = result + return response, shutdown +} + +func (s *server) dispatch(method string, params jsonObject) (any, bool, error) { + if _, unsupported := allVhostsUnsupportedMethods[method]; unsupported && allVhostsRequested(params) { + return nil, false, errors.New("all_vhosts is only supported for list operations") + } + switch method { + case "handshake": + return handshakeResult{protocolVersion, agentProtocolVersion, capabilities}, false, nil + case "connect": + result, err := s.connect(params) + return result, false, err + case "test_connection": + result, err := s.testConnection(params) + return result, false, err + case "disconnect": + s.closeClients() + return okResult(), false, nil + case "shutdown": + s.closeClients() + return okResult(), true, nil + case "mq_list_topics": + result, err := s.listTopics(params) + return result, false, err + case "mq_create_topic": + result, err := s.createTopic(params) + return result, false, err + case "mq_delete_topic": + result, err := s.deleteTopic(params) + return result, false, err + case "mq_get_topic_stats": + result, err := s.getTopicStats(params) + return result, false, err + case "mq_get_topic_config": + result, err := s.getTopicConfig(params) + return result, false, err + case "mq_alter_topic_config": + return nil, false, errors.New("RabbitMQ queue arguments are immutable after declaration; delete and re-declare the queue to change them") + case "mq_purge_queue": + result, err := s.purgeQueue(params) + return result, false, err + case "mq_list_consumers": + result, err := s.listConsumers(params) + return result, false, err + case "mq_list_namespaces": + result, err := s.listNamespaces(params) + return result, false, err + case "mq_create_namespace": + result, err := s.createNamespace(params) + return result, false, err + case "mq_delete_namespace": + result, err := s.deleteNamespace(params) + return result, false, err + case "mq_list_exchanges": + result, err := s.listExchanges(params) + return result, false, err + case "mq_create_exchange": + result, err := s.createExchange(params) + return result, false, err + case "mq_delete_exchange": + result, err := s.deleteExchange(params) + return result, false, err + case "mq_list_bindings": + result, err := s.listBindings(params) + return result, false, err + case "mq_bind": + result, err := s.bind(params) + return result, false, err + case "mq_unbind": + result, err := s.unbind(params) + return result, false, err + case "mq_list_connections": + result, err := s.listClientConnections(params) + return result, false, err + case "mq_list_channels": + result, err := s.listClientChannels(params) + return result, false, err + case "mq_close_connection": + result, err := s.closeClientConnection(params) + return result, false, err + case "mq_list_users": + result, err := s.listUsers(params) + return result, false, err + case "mq_create_user": + result, err := s.createUser(params) + return result, false, err + case "mq_delete_user": + result, err := s.deleteUser(params) + return result, false, err + case "mq_list_permissions": + result, err := s.listPermissions(params) + return result, false, err + case "mq_grant_permission": + result, err := s.grantPermission(params) + return result, false, err + case "mq_revoke_permission": + result, err := s.revokePermission(params) + return result, false, err + case "mq_list_policies": + result, err := s.listPolicies(params) + return result, false, err + case "mq_set_policy": + result, err := s.setPolicy(params) + return result, false, err + case "mq_delete_policy": + result, err := s.deletePolicy(params) + return result, false, err + case "mq_peek_messages": + result, err := s.peekMessages(params) + return result, false, err + case "mq_send_message": + result, err := s.sendMessage(params) + return result, false, err + case "mq_describe_cluster": + result, err := s.describeCluster(params) + return result, false, err + case "mq_overview": + result, err := s.getOverview(params) + return result, false, err + case "mq_list_nodes": + result, err := s.listNodes(params) + return result, false, err + default: + return nil, false, fmt.Errorf("Unknown method: %s", method) + } +} + +func (s *server) connect(params jsonObject) (any, error) { + config := connectionObject(params) + nextConnection, err := openConnection(config) + if err != nil { + return nil, err + } + nextChannel, err := nextConnection.Channel() + if err != nil { + closeConnection(nextConnection) + return nil, err + } + s.closeClients() + s.connection = nextConnection + s.channel = nextChannel + s.cachedConnection = deepCopyObject(config) + return okResult(), nil +} + +func (s *server) testConnection(params jsonObject) (any, error) { + config := connectionObject(params) + connection, err := openConnection(config) + if err != nil { + return nil, err + } + defer closeConnection(connection) + version := serverString(connection.Properties, "version") + return jsonObject{ + "ok": true, + "product": serverString(connection.Properties, "product"), + "version": version, + "serverVersion": version, + "clusterName": serverString(connection.Properties, "cluster_name"), + "platform": serverString(connection.Properties, "platform"), + }, nil +} + +func (s *server) closeClients() { + for key, client := range s.vhostClients { + client.close() + delete(s.vhostClients, key) + } + closeChannel(s.channel) + s.channel = nil + closeConnection(s.connection) + s.connection = nil + s.cachedConnection = nil +} + +func (client *vhostClient) close() { + if client == nil { + return + } + closeChannel(client.channel) + closeConnection(client.connection) +} + +func (client *vhostClient) isOpen() bool { + return client != nil && client.connection != nil && !client.connection.IsClosed() && client.channel != nil && !client.channel.IsClosed() +} + +func closeChannel(channel *amqp.Channel) { + if channel != nil { + _ = channel.Close() + } +} + +func closeConnection(connection *amqp.Connection) { + if connection != nil { + _ = connection.Close() + } +} + +func openConnection(config jsonObject) (*amqp.Connection, error) { + addresses, err := resolveAddresses(config) + if err != nil { + return nil, err + } + var lastError error + for _, endpoint := range addresses { + connection, dialError := dialAddress(config, endpoint) + if dialError == nil { + return connection, nil + } + lastError = dialError + } + if lastError == nil { + return nil, errors.New("addresses is required") + } + return nil, lastError +} + +func dialAddress(config jsonObject, endpoint address) (*amqp.Connection, error) { + properties := objectOrNil(config, "properties") + connectionTimeout := durationMilliseconds(config, "request_timeout_ms", defaultRequestTimeout) + if configured, ok := integerProperty(properties, "connection_timeout_ms"); ok { + connectionTimeout = time.Duration(configured) * time.Millisecond + } + handshakeTimeout := defaultHandshakeTimeout + if configured, ok := integerProperty(properties, "handshake_timeout_ms"); ok { + handshakeTimeout = time.Duration(configured) * time.Millisecond + } + heartbeat := int(defaultHeartbeat / time.Second) + if configured, ok := integerProperty(properties, "requested_heartbeat"); ok { + heartbeat = configured + } + scheme := "amqp" + if amqpTLSEnabled(config) { + scheme = "amqps" + } + uri := url.URL{Scheme: scheme, Host: net.JoinHostPort(endpoint.Host, strconv.Itoa(endpoint.Port)), Path: "/"} + query := uri.Query() + query.Set("heartbeat", strconv.Itoa(heartbeat)) + uri.RawQuery = query.Encode() + amqpConfig := amqp.Config{ + SASL: []amqp.Authentication{&amqp.PlainAuth{ + Username: credentialOrGuest(config, "username"), + Password: credentialOrGuest(config, "password"), + }}, + Vhost: stringOrDefault(config, "virtual_host", "/"), + ChannelMax: defaultChannelMax, + Heartbeat: time.Duration(heartbeat) * time.Second, + Dial: func(network, target string) (net.Conn, error) { + dialer := net.Dialer{Timeout: connectionTimeout} + connection, err := dialer.Dial(network, target) + if err != nil { + return nil, err + } + if err := connection.SetDeadline(time.Now().Add(handshakeTimeout)); err != nil { + _ = connection.Close() + return nil, err + } + return connection, nil + }, + } + if scheme == "amqps" { + amqpConfig.TLSClientConfig = &tls.Config{ + ServerName: endpoint.Host, + InsecureSkipVerify: tlsSkipVerify(config), + } + } + return amqp.DialConfig(uri.String(), amqpConfig) +} + +func (s *server) channelFor(params jsonObject) (*amqp.Channel, error) { + defaultVhost := "/" + if s.cachedConnection != nil { + defaultVhost = stringOrDefault(s.cachedConnection, "virtual_host", "/") + } + vhost := effectiveVhost(params, s.cachedConnection) + if vhost == defaultVhost { + return s.primaryChannel() + } + if s.cachedConnection == nil { + return nil, errors.New("Not connected. Call connect first.") + } + if client := s.vhostClients[vhost]; client != nil && client.isOpen() { + return client.channel, nil + } else if client != nil { + client.close() + delete(s.vhostClients, vhost) + } + config := deepCopyObject(s.cachedConnection) + config["virtual_host"] = vhost + connection, err := openConnection(config) + if err != nil { + return nil, err + } + channel, err := connection.Channel() + if err != nil { + closeConnection(connection) + return nil, err + } + s.vhostClients[vhost] = &vhostClient{connection: connection, channel: channel} + return channel, nil +} + +func (s *server) primaryChannel() (*amqp.Channel, error) { + if s.connection != nil && !s.connection.IsClosed() && !needsNewChannel(s.channel) { + return s.channel, nil + } + if s.connection == nil || s.connection.IsClosed() { + if s.cachedConnection == nil { + return nil, errors.New("Not connected. Call connect first.") + } + closeConnection(s.connection) + connection, err := openConnection(s.cachedConnection) + if err != nil { + return nil, err + } + s.connection = connection + } + closeChannel(s.channel) + channel, err := s.connection.Channel() + if err != nil { + return nil, err + } + s.channel = channel + return channel, nil +} + +func needsNewChannel(channel *amqp.Channel) bool { + return channel == nil || channel.IsClosed() +} + +func resolveAddresses(config jsonObject) ([]address, error) { + addresses := strings.TrimSpace(stringOrEmpty(config, "addresses")) + if addresses == "" { + addresses = strings.TrimSpace(stringOrEmpty(config, "host")) + } + if addresses == "" { + return nil, errors.New("addresses is required") + } + return parseAddresses(addresses, intOrDefault(config, "port", defaultAMQPPort)) +} + +func parseAddresses(value string, defaultPort int) ([]address, error) { + result := make([]address, 0) + for _, part := range strings.Split(value, ",") { + trimmed := strings.TrimSpace(part) + if trimmed == "" { + continue + } + host := trimmed + port := defaultPort + if parsedHost, parsedPort, err := net.SplitHostPort(trimmed); err == nil { + host = parsedHost + parsed, parseError := strconv.Atoi(parsedPort) + if parseError != nil { + return nil, parseError + } + port = parsed + } else if colon := strings.LastIndex(trimmed, ":"); colon > 0 && colon < len(trimmed)-1 && strings.Count(trimmed, ":") == 1 { + parsed, parseError := strconv.Atoi(trimmed[colon+1:]) + if parseError != nil { + return nil, parseError + } + host = trimmed[:colon] + port = parsed + } else { + host = strings.TrimPrefix(strings.TrimSuffix(trimmed, "]"), "[") + } + result = append(result, address{Host: host, Port: port}) + } + if len(result) == 0 { + return nil, errors.New("addresses is required") + } + return result, nil +} + +func amqpTLSEnabled(config jsonObject) bool { + _, nestedTLS := config["tls"].(map[string]any) + if !nestedTLS { + _, nestedTLS = config["tls"].(jsonObject) + } + return nestedTLS || boolOrDefault(config, "tls_skip_verify", false) || boolProperty(config, "ssl") || boolProperty(config, "tls") +} + +func tlsSkipVerify(config jsonObject) bool { + if boolOrDefault(config, "tls_skip_verify", false) { + return true + } + tlsConfig := objectOrNil(config, "tls") + return boolOrDefault(tlsConfig, "skip_verify", false) +} diff --git a/agents/drivers/rabbitmq/management.go b/agents/drivers/rabbitmq/management.go new file mode 100644 index 000000000..d78056afc --- /dev/null +++ b/agents/drivers/rabbitmq/management.go @@ -0,0 +1,248 @@ +package main + +import ( + "bytes" + "context" + "crypto/tls" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "strconv" + "strings" + "time" +) + +const ( + defaultManagementPort = 15672 + defaultManagementTLSPort = 15671 + managementPageSize = 100 + managementConnectTimeout = 10 * time.Second + managementRequestTimeout = 20 * time.Second +) + +type managementStatusError struct { + status int + method string + path string +} + +func (err *managementStatusError) Error() string { + return managementErrorMessage(err.status, err.method, err.path) +} + +func managementGet(connection jsonObject, path string) (any, error) { + return managementRequest(connection, http.MethodGet, path, nil) +} + +func managementSend(connection jsonObject, method, path string, body jsonObject) (any, error) { + return managementRequest(connection, method, path, body) +} + +func managementRequest(connection jsonObject, method, path string, body jsonObject) (any, error) { + baseURLs, err := managementBaseURLs(connection) + if err != nil { + return nil, err + } + var lastConnectionError error + for _, baseURL := range baseURLs { + result, requestError := managementRequestOnce(baseURL, connection, method, path, body) + if requestError == nil { + return result, nil + } + var statusError *managementStatusError + if errors.As(requestError, &statusError) { + return nil, requestError + } + var networkError net.Error + if errors.As(requestError, &networkError) { + lastConnectionError = requestError + continue + } + return nil, requestError + } + if lastConnectionError != nil { + return nil, lastConnectionError + } + return nil, errors.New("No management API endpoint candidates") +} + +func managementRequestOnce(baseURL string, connection jsonObject, method, path string, body jsonObject) (any, error) { + var requestBody io.Reader + if body != nil { + encoded, err := json.Marshal(body) + if err != nil { + return nil, err + } + requestBody = bytes.NewReader(encoded) + } + ctx, cancel := context.WithTimeout(context.Background(), managementRequestTimeout) + defer cancel() + request, err := http.NewRequestWithContext(ctx, method, baseURL+path, requestBody) + if err != nil { + return nil, err + } + request.Header.Set("Authorization", basicAuthHeader( + credentialOrGuest(connection, "username"), credentialOrGuest(connection, "password"))) + if body != nil { + request.Header.Set("Content-Type", "application/json") + } + transport := &http.Transport{ + DialContext: (&net.Dialer{Timeout: managementConnectTimeout}).DialContext, + TLSHandshakeTimeout: managementConnectTimeout, + ResponseHeaderTimeout: managementConnectTimeout, + TLSClientConfig: &tls.Config{InsecureSkipVerify: tlsSkipVerify(connection)}, + } + defer transport.CloseIdleConnections() + client := &http.Client{Transport: transport} + response, err := client.Do(request) + if err != nil { + return nil, err + } + defer response.Body.Close() + if response.StatusCode < 200 || response.StatusCode >= 300 { + return nil, &managementStatusError{status: response.StatusCode, method: method, path: path} + } + if response.StatusCode == http.StatusNoContent { + return nil, nil + } + data, err := io.ReadAll(response.Body) + if err != nil { + return nil, err + } + if strings.TrimSpace(string(data)) == "" { + return nil, nil + } + var result any + if err := decodeJSON(data, &result); err != nil { + return nil, err + } + return result, nil +} + +func managementGetAll(connection jsonObject, path string) ([]any, error) { + all := make([]any, 0) + for page := 1; ; page++ { + separator := "?" + if strings.Contains(path, "?") { + separator = "&" + } + response, err := managementGet(connection, + path+separator+"page="+strconv.Itoa(page)+"&page_size="+strconv.Itoa(managementPageSize)) + if err != nil { + return nil, err + } + switch typed := response.(type) { + case []any: + return append(all, typed...), nil + case map[string]any: + items, exists := typed["items"] + if !exists { + return nil, fmt.Errorf("Unexpected management API response for list endpoint %s", path) + } + if array, ok := items.([]any); ok { + all = append(all, array...) + } + pageCount := integerOrNull(jsonObject(typed), "page_count") + if pageCount == nil || page >= *pageCount { + return all, nil + } + default: + return nil, fmt.Errorf("Unexpected management API response for list endpoint %s", path) + } + } +} + +func managementBaseURLs(connection jsonObject) ([]string, error) { + if explicit := stringOrNull(connection, "management_url"); explicit != nil && strings.TrimSpace(*explicit) != "" { + return []string{normalizeManagementURL(*explicit)}, nil + } + tlsEnabled := managementTLS(connection) + port := managementPort(connection, tlsEnabled) + addresses, err := resolveAddresses(connection) + if err != nil { + return nil, err + } + baseURLs := make([]string, 0, len(addresses)) + for _, endpoint := range addresses { + baseURLs = append(baseURLs, managementBaseURL(endpoint.Host, port, tlsEnabled)) + } + return baseURLs, nil +} + +func managementBaseURL(host string, port int, tlsEnabled bool) string { + scheme := "http" + if tlsEnabled { + scheme = "https" + } + return scheme + "://" + net.JoinHostPort(host, strconv.Itoa(port)) +} + +func normalizeManagementURL(value string) string { + return strings.TrimRight(strings.TrimSpace(value), "/") +} + +func managementTLS(connection jsonObject) bool { + return objectOrNil(connection, "tls") != nil || boolProperty(connection, "ssl") || boolProperty(connection, "tls") +} + +func managementPort(connection jsonObject, tlsEnabled bool) int { + if configured, ok := integerProperty(objectOrNil(connection, "properties"), "management_port"); ok { + return configured + } + if tlsEnabled { + return defaultManagementTLSPort + } + return defaultManagementPort +} + +func credentialOrGuest(connection jsonObject, key string) string { + value := stringOrNull(connection, key) + if value == nil || strings.TrimSpace(*value) == "" { + return "guest" + } + return *value +} + +func basicAuthHeader(username, password string) string { + return "Basic " + base64.StdEncoding.EncodeToString([]byte(username+":"+password)) +} + +func managementErrorMessage(status int, method, path string) string { + base := fmt.Sprintf("RabbitMQ management API returned HTTP %d for %s %s.", status, method, path) + if status == http.StatusUnauthorized || status == http.StatusForbidden { + return base + " Hint: check the username/password and that the user has a management permission tag (management, policymaker, monitoring, or administrator)." + } + return base + " The rabbitmq_management plugin must be enabled for this operation." +} + +func urlEncodeVhost(value string) string { + return javaFormPathEscape(value) +} + +func urlEncodePathSegment(value string) string { + return javaFormPathEscape(value) +} + +func urlEncodeName(value string) string { + return urlEncodePathSegment(value) +} + +func javaFormPathEscape(value string) string { + const hex = "0123456789ABCDEF" + var builder strings.Builder + for _, current := range []byte(value) { + if (current >= 'a' && current <= 'z') || (current >= 'A' && current <= 'Z') || + (current >= '0' && current <= '9') || current == '-' || current == '_' || current == '.' || current == '*' { + builder.WriteByte(current) + continue + } + builder.WriteByte('%') + builder.WriteByte(hex[current>>4]) + builder.WriteByte(hex[current&15]) + } + return builder.String() +} diff --git a/agents/drivers/rabbitmq/management_test.go b/agents/drivers/rabbitmq/management_test.go new file mode 100644 index 000000000..eba67c121 --- /dev/null +++ b/agents/drivers/rabbitmq/management_test.go @@ -0,0 +1,221 @@ +package main + +import ( + "encoding/json" + "io" + "net" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "sync" + "testing" +) + +func TestManagementBaseURLs(t *testing.T) { + explicit, err := managementBaseURLs(mustObject(t, `{"management_url":" https://proxy:8443/rmq/ "}`)) + if err != nil || len(explicit) != 1 || explicit[0] != "https://proxy:8443/rmq" { + t.Fatalf("unexpected explicit URLs %#v, %v", explicit, err) + } + withoutAddresses, err := managementBaseURLs(mustObject(t, `{"management_url":"http://mgmt:15672"}`)) + if err != nil || withoutAddresses[0] != "http://mgmt:15672" { + t.Fatalf("unexpected URL %#v, %v", withoutAddresses, err) + } + derived, err := managementBaseURLs(mustObject(t, `{"addresses":"mq1:5672,mq2:5673"}`)) + if err != nil || len(derived) != 2 || derived[0] != "http://mq1:15672" || derived[1] != "http://mq2:15672" { + t.Fatalf("unexpected derived URLs %#v, %v", derived, err) + } + tlsDerived, err := managementBaseURLs(mustObject(t, `{"addresses":"mq1","tls":{}}`)) + if err != nil || tlsDerived[0] != "https://mq1:15671" { + t.Fatalf("unexpected TLS URLs %#v, %v", tlsDerived, err) + } + skipVerify, err := managementBaseURLs(mustObject(t, `{"addresses":"mq1","tls_skip_verify":true}`)) + if err != nil || skipVerify[0] != "http://mq1:15672" { + t.Fatalf("unexpected skip-verify URLs %#v, %v", skipVerify, err) + } +} + +func TestManagementErrorMessages(t *testing.T) { + for _, status := range []int{401, 403} { + message := managementErrorMessage(status, http.MethodGet, "/api/queues") + if !strings.Contains(message, "management permission tag") || strings.Contains(message, "plugin must be enabled") { + t.Fatalf("unexpected message %q", message) + } + } + message := managementErrorMessage(404, http.MethodGet, "/api/queues/%2F/gone") + if !strings.Contains(message, "plugin must be enabled") || strings.Contains(message, "management permission tag") { + t.Fatalf("unexpected message %q", message) + } +} + +func TestManagementRequestSurfacesCredentialError(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.WriteHeader(http.StatusUnauthorized) + })) + defer server.Close() + connection := jsonObject{"management_url": server.URL} + _, err := managementGet(connection, "/api/queues") + if err == nil || !strings.Contains(err.Error(), "HTTP 401") || !strings.Contains(err.Error(), "management permission tag") { + t.Fatalf("unexpected error %v", err) + } +} + +func TestManagementGetAllPagination(t *testing.T) { + var mutex sync.Mutex + requestedPages := make([]int, 0) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + page, _ := strconv.Atoi(request.URL.Query().Get("page")) + mutex.Lock() + requestedPages = append(requestedPages, page) + mutex.Unlock() + writer.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(writer, `{"items":[{"name":"q`+strconv.Itoa(page)+`"}],"page":`+strconv.Itoa(page)+`,"page_count":3,"total_count":3}`) + })) + defer server.Close() + items, err := managementGetAll(jsonObject{"management_url": server.URL}, "/api/queues") + if err != nil { + t.Fatal(err) + } + if len(items) != 3 || len(requestedPages) != 3 || requestedPages[0] != 1 || requestedPages[2] != 3 { + t.Fatalf("unexpected items %#v pages %#v", items, requestedPages) + } +} + +func TestManagementGetAllPlainArray(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(writer, `[{"name":"guest"}]`) + })) + defer server.Close() + items, err := managementGetAll(jsonObject{"management_url": server.URL}, "/api/users") + if err != nil || len(items) != 1 { + t.Fatalf("unexpected items %#v, %v", items, err) + } +} + +func TestManagementRequestFailsOverConnectionErrors(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.2:0") + if err != nil { + t.Skipf("secondary loopback address unavailable: %v", err) + } + port := listener.Addr().(*net.TCPAddr).Port + server := &http.Server{Handler: http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(writer, `[]`) + })} + defer server.Close() + go server.Serve(listener) + connection := mustObject(t, `{"addresses":"127.0.0.1,127.0.0.2","properties":{"management_port":`+strconv.Itoa(port)+`}}`) + response, err := managementGet(connection, "/api/queues") + if err != nil { + t.Fatal(err) + } + if _, ok := response.([]any); !ok { + t.Fatalf("unexpected response %#v", response) + } +} + +func TestManagementHTTPErrorDoesNotFailOver(t *testing.T) { + first, second, port := pairedLoopbackServers(t, + http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.WriteHeader(http.StatusNotFound) + }), + http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(writer, `[]`) + }), + ) + defer first.Close() + defer second.Close() + connection := mustObject(t, `{"addresses":"127.0.0.1,127.0.0.2","properties":{"management_port":`+strconv.Itoa(port)+`}}`) + _, err := managementGet(connection, "/api/queues") + if err == nil || !strings.Contains(err.Error(), "HTTP 404") { + t.Fatalf("unexpected error %v", err) + } +} + +func TestManagementURLPathPrefix(t *testing.T) { + requestedPath := "" + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + requestedPath = request.URL.Path + writer.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(writer, `[]`) + })) + defer server.Close() + _, err := managementGet(jsonObject{"management_url": server.URL + "/rmq/"}, "/api/queues") + if err != nil { + t.Fatal(err) + } + if requestedPath != "/rmq/api/queues" { + t.Fatalf("got path %q", requestedPath) + } +} + +func TestPolicyManagementOperations(t *testing.T) { + type capturedRequest struct { + Method string + Path string + Body jsonObject + } + requests := make([]capturedRequest, 0) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + captured := capturedRequest{Method: request.Method, Path: request.URL.Path} + if request.Body != nil { + data, _ := io.ReadAll(request.Body) + if len(data) > 0 { + _ = decodeJSON(data, &captured.Body) + } + } + requests = append(requests, captured) + writer.Header().Set("Content-Type", "application/json") + switch request.Method { + case http.MethodGet: + _, _ = io.WriteString(writer, `[{"name":"ha","vhost":"/","pattern":"^ha","apply-to":"queues","priority":0,"definition":{"ha-mode":"all"}}]`) + default: + writer.WriteHeader(http.StatusNoContent) + } + })) + defer server.Close() + service := newServer() + service.cachedConnection = jsonObject{"management_url": server.URL, "username": "guest", "password": "guest"} + listed, err := service.listPolicies(jsonObject{"virtual_host": "/"}) + if err != nil { + t.Fatal(err) + } + if len(listed.(jsonObject)["policies"].([]jsonObject)) != 1 { + t.Fatalf("unexpected policies %#v", listed) + } + _, err = service.setPolicy(mustObject(t, `{"virtual_host":"/","name":"ha","pattern":"^ha","definition":{"ha-mode":"all"}}`)) + if err != nil { + t.Fatal(err) + } + _, err = service.deletePolicy(mustObject(t, `{"virtual_host":"/","name":"ha"}`)) + if err != nil { + t.Fatal(err) + } + if len(requests) != 3 || requests[1].Method != http.MethodPut || requests[2].Method != http.MethodDelete { + t.Fatalf("unexpected requests %#v", requests) + } + if requests[1].Body["apply-to"] != "queues" || requests[1].Body["priority"] != json.Number("0") { + t.Fatalf("unexpected body %#v", requests[1].Body) + } +} + +func pairedLoopbackServers(t *testing.T, firstHandler, secondHandler http.Handler) (*http.Server, *http.Server, int) { + t.Helper() + firstListener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + port := firstListener.Addr().(*net.TCPAddr).Port + secondListener, err := net.Listen("tcp", "127.0.0.2:"+strconv.Itoa(port)) + if err != nil { + firstListener.Close() + t.Skipf("secondary loopback address unavailable: %v", err) + } + first := &http.Server{Handler: firstHandler} + second := &http.Server{Handler: secondHandler} + go first.Serve(firstListener) + go second.Serve(secondListener) + return first, second, port +} diff --git a/agents/drivers/rabbitmq/mapping_test.go b/agents/drivers/rabbitmq/mapping_test.go new file mode 100644 index 000000000..61885a952 --- /dev/null +++ b/agents/drivers/rabbitmq/mapping_test.go @@ -0,0 +1,126 @@ +package main + +import "testing" + +func TestConsumersFromQueueInfo(t *testing.T) { + info := mustObject(t, `{ + "consumer_details": [ + {"consumer_tag":"ctag","active":true,"ack_required":true,"prefetch_count":25,"channel_details":{"name":"conn (1)"}}, + "ignored" + ] + }`) + consumers := consumersFromQueueInfo(info) + if len(consumers) != 1 || consumers[0]["name"] != "conn (1)" || consumers[0]["tag"] != "ctag" || consumers[0]["prefetch"] != 25 { + t.Fatalf("unexpected consumers %#v", consumers) + } + if got := consumersFromQueueInfo(jsonObject{}); len(got) != 0 { + t.Fatalf("unexpected consumers %#v", got) + } +} + +func TestExchangeAndBindingMappings(t *testing.T) { + defaultExchange := exchangeInfoFromJSON(mustObject(t, `{"name":"","type":"","durable":true,"auto_delete":false,"internal":false}`)) + if defaultExchange["type"] != "default" || defaultExchange["durable"] != true { + t.Fatalf("unexpected exchange %#v", defaultExchange) + } + topicExchange := exchangeInfoFromJSON(mustObject(t, `{"name":"events","type":"topic","durable":true,"auto_delete":true,"internal":true}`)) + if topicExchange["type"] != "topic" || topicExchange["autoDelete"] != true || topicExchange["internal"] != true { + t.Fatalf("unexpected exchange %#v", topicExchange) + } + binding := bindingInfoFromJSON(mustObject(t, `{ + "source":"events","destination":"orders","destination_type":"queue","routing_key":"orders.*", + "arguments":{"x-priority":5,"alternate":true,"ignored":null} + }`)) + if binding["destinationType"] != "queue" || binding["routingKey"] != "orders.*" { + t.Fatalf("unexpected binding %#v", binding) + } + arguments := binding["arguments"].(jsonObject) + if arguments["x-priority"] != int64(5) || arguments["alternate"] != true { + t.Fatalf("unexpected arguments %#v", arguments) + } + withoutArguments := bindingInfoFromJSON(mustObject(t, `{"source":"e","destination":"q","destination_type":"queue","routing_key":"","arguments":{}}`)) + if _, exists := withoutArguments["arguments"]; exists { + t.Fatalf("unexpected arguments %#v", withoutArguments) + } +} + +func TestConnectionAndChannelMappings(t *testing.T) { + connection := clientConnectionInfoFromJSON(mustObject(t, `{ + "name":"127.0.0.1:1 -> 127.0.0.1:5672","user":"dbx","peer_host":"127.0.0.1","peer_port":1234, + "state":"running","channels":2,"recv_oct_details":{"rate":12.5},"send_oct_details":{"rate":8.25},"connected_at":1700000000000 + }`)) + if connection["recvRate"] != 12.5 || connection["sendRate"] != 8.25 || connection["connectedAt"] != int64(1700000000000) { + t.Fatalf("unexpected connection %#v", connection) + } + minimal := clientConnectionInfoFromJSON(mustObject(t, `{"name":"conn"}`)) + for _, key := range []string{"recvRate", "sendRate", "connectedAt"} { + if _, exists := minimal[key]; exists { + t.Fatalf("unexpected %s in %#v", key, minimal) + } + } + channel := channelInfoFromJSON(mustObject(t, `{ + "name":"conn (1)","connection_details":{"name":"conn"},"state":"running", + "prefetch_count":10,"messages_unacknowledged":4,"consumer_count":2 + }`)) + if channel["connectionName"] != "conn" || channel["messagesUnacked"] != int64(4) || channel["consumerCount"] != int64(2) { + t.Fatalf("unexpected channel %#v", channel) + } + if !channelMatchesConnection(channel, "conn") || !channelMatchesConnection(jsonObject{"name": "other (1)"}, "other") { + t.Fatal("connection matching failed") + } + if channelMatchesConnection(channel, "missing") { + t.Fatal("unexpected connection match") + } +} + +func TestUserPermissionPolicyMappings(t *testing.T) { + user := userInfoFromJSON(mustObject(t, `{"name":"admin","tags":"administrator, management"}`)) + if user["name"] != "admin" || len(user["tags"].([]string)) != 2 { + t.Fatalf("unexpected user %#v", user) + } + permission := permissionInfoFromJSON(mustObject(t, `{"user":"dbx","vhost":"/","configure":".*","write":"^orders","read":".*"}`)) + if permission["write"] != "^orders" || permission["vhost"] != "/" { + t.Fatalf("unexpected permission %#v", permission) + } + policy := policyInfoFromJSON(mustObject(t, `{ + "name":"ha","vhost":"/","pattern":"^ha","apply-to":"queues","priority":5, + "definition":{"ha-mode":"all","ha-sync-mode":"automatic","expires":60000,"ignored":null} + }`)) + if policy["applyTo"] != "queues" || policy["priority"] != int64(5) { + t.Fatalf("unexpected policy %#v", policy) + } + definition := policy["definition"].(jsonObject) + if definition["expires"] != int64(60000) || definition["ha-mode"] != "all" { + t.Fatalf("unexpected definition %#v", definition) + } +} + +func TestOverviewAndNodeMappings(t *testing.T) { + overview := overviewInfoFromJSON(mustObject(t, `{ + "queue_totals":{"messages_ready":10,"messages_unacknowledged":2}, + "message_stats":{"publish_details":{"rate":1.5},"deliver_get_details":{"rate":2.5},"ack_details":{"rate":3.5}}, + "object_totals":{"queues":4,"exchanges":5,"connections":6,"channels":7,"consumers":8} + }`)) + if overview["messagesReady"] != int64(10) || overview["publishRate"] != 1.5 || overview["totalConsumers"] != int64(8) { + t.Fatalf("unexpected overview %#v", overview) + } + minimal := overviewInfoFromJSON(jsonObject{}) + if len(minimal) != 0 { + t.Fatalf("unexpected overview %#v", minimal) + } + node := nodeInfoFromJSON(mustObject(t, `{ + "name":"rabbit@node","running":true,"mem_used":100,"mem_limit":200,"disk_free":300, + "fd_used":4,"fd_total":5,"sockets_used":6,"sockets_total":7,"uptime":8000 + }`)) + if node["running"] != true || node["memUsed"] != int64(100) || node["uptimeMs"] != int64(8000) { + t.Fatalf("unexpected node %#v", node) + } +} + +func TestAttachVhost(t *testing.T) { + info := jsonObject{"name": "q"} + attachVhost(info, mustObject(t, `{"vhost":"orders"}`)) + if info["vhost"] != "orders" { + t.Fatalf("unexpected info %#v", info) + } +} diff --git a/agents/drivers/rabbitmq/operations.go b/agents/drivers/rabbitmq/operations.go new file mode 100644 index 000000000..51854f923 --- /dev/null +++ b/agents/drivers/rabbitmq/operations.go @@ -0,0 +1,1345 @@ +package main + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "net/http" + "os" + "sort" + "strings" + "unicode/utf8" + + amqp "github.com/rabbitmq/amqp091-go" +) + +const ( + maxPeekMessages = 10000 + defaultPermissionPattern = ".*" +) + +var exchangeTypes = map[string]struct{}{ + "direct": {}, "fanout": {}, "topic": {}, "headers": {}, +} + +func (s *server) listTopics(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + allVhosts := allVhostsRequested(params) + queues, err := managementGetAll(connection, managementListPath(params, connection, "queues")) + if err != nil { + return nil, err + } + topics := make([]jsonObject, 0, len(queues)) + for _, value := range queues { + queue, ok := value.(map[string]any) + if !ok { + continue + } + info := jsonObject{ + "name": stringOrEmpty(jsonObject(queue), "name"), + "durable": boolOrDefault(jsonObject(queue), "durable", false), + "autoDelete": boolOrDefault(jsonObject(queue), "auto_delete", false), + "state": stringOrEmpty(jsonObject(queue), "state"), + "messages": longOrDefault(jsonObject(queue), "messages", 0), + "consumers": longOrDefault(jsonObject(queue), "consumers", 0), + } + if allVhosts { + attachVhost(info, jsonObject(queue)) + } + topics = append(topics, info) + } + sort.SliceStable(topics, func(left, right int) bool { + return stringOrEmpty(topics[left], "name") < stringOrEmpty(topics[right], "name") + }) + return jsonObject{"topics": topics}, nil +} + +func (s *server) createTopic(params jsonObject) (any, error) { + channel, err := s.channelFor(params) + if err != nil { + return nil, err + } + name, err := queueName(params) + if err != nil { + return nil, err + } + arguments := amqp.Table{} + if configs := objectOrNil(params, "configs"); configs != nil { + for key, value := range configs { + if converted := argumentValue(value); converted != nil { + arguments[key] = converted + } + } + } + _, err = channel.QueueDeclare(name, boolOrDefault(params, "durable", true), false, false, false, arguments) + if err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) deleteTopic(params jsonObject) (any, error) { + channel, err := s.channelFor(params) + if err != nil { + return nil, err + } + name, err := queueName(params) + if err != nil { + return nil, err + } + if _, err := channel.QueueDelete(name, false, false, false); err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) getTopicStats(params jsonObject) (any, error) { + name, err := queueName(params) + if err != nil { + return nil, err + } + if connection := s.currentConnectionConfig(params); connection != nil { + vhost := effectiveVhost(params, connection) + queue, managementError := managementGet(connection, + "/api/queues/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(name)) + if managementError == nil { + if info, ok := queue.(map[string]any); ok { + messages := longOrDefault(jsonObject(info), "messages", 0) + return jsonObject{ + "name": name, + "messageCount": messages, + "consumerCount": longOrDefault(jsonObject(info), "consumers", 0), + "totalMessages": messages, + }, nil + } + } else { + fmt.Fprintln(os.Stderr, "Management API unavailable for queue stats, falling back to passive declare: "+managementError.Error()) + } + } + channel, err := s.channelFor(params) + if err != nil { + return nil, err + } + queue, err := channel.QueueDeclarePassive(name, false, false, false, false, nil) + if err != nil { + return nil, err + } + return jsonObject{ + "name": name, + "messageCount": queue.Messages, + "consumerCount": queue.Consumers, + "totalMessages": queue.Messages, + }, nil +} + +func (s *server) getTopicConfig(params jsonObject) (any, error) { + channel, err := s.channelFor(params) + if err != nil { + return nil, err + } + name, err := queueName(params) + if err != nil { + return nil, err + } + configs := jsonObject{} + if connection := s.currentConnectionConfig(params); connection != nil { + vhost := effectiveVhost(params, connection) + queue, managementError := managementGet(connection, + "/api/queues/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(name)) + if managementError == nil { + if info, ok := queue.(map[string]any); ok { + object := jsonObject(info) + configs["durable"] = boolOrDefault(object, "durable", false) + configs["auto_delete"] = boolOrDefault(object, "auto_delete", false) + configs["exclusive"] = boolOrDefault(object, "exclusive", false) + if arguments := objectOrNil(object, "arguments"); arguments != nil { + for key, value := range arguments { + if value == nil { + configs[key] = nil + } else { + configs[key] = fmt.Sprint(value) + } + } + } + } + } else { + fmt.Fprintln(os.Stderr, "Management API unavailable for queue config: "+managementError.Error()) + } + } + if len(configs) == 0 { + if _, err := channel.QueueDeclarePassive(name, false, false, false, false, nil); err != nil { + return nil, err + } + } + return jsonObject{"configs": configs}, nil +} + +func (s *server) purgeQueue(params jsonObject) (any, error) { + name, err := queueName(params) + if err != nil { + return nil, err + } + channel, err := s.channelFor(params) + if err != nil { + return nil, err + } + purged, err := channel.QueuePurge(name, false) + if err != nil { + return nil, err + } + return jsonObject{"ok": true, "purged": purged}, nil +} + +func (s *server) listConsumers(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + name, err := queueName(params) + if err != nil { + return nil, err + } + vhost := effectiveVhost(params, connection) + queue, err := managementGet(connection, + "/api/queues/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(name)) + if err != nil { + return nil, err + } + info, ok := queue.(map[string]any) + if !ok { + return nil, errors.New("Unexpected management API response for queue details") + } + return jsonObject{"consumers": consumersFromQueueInfo(jsonObject(info))}, nil +} + +func consumersFromQueueInfo(info jsonObject) []jsonObject { + details := arrayOrNil(info, "consumer_details") + consumers := make([]jsonObject, 0, len(details)) + for _, value := range details { + consumerMap, ok := value.(map[string]any) + if !ok { + continue + } + consumer := jsonObject(consumerMap) + channelName := "" + if channelDetails := objectOrNil(consumer, "channel_details"); channelDetails != nil { + channelName = stringOrEmpty(channelDetails, "name") + } + entry := jsonObject{ + "name": channelName, + "tag": stringOrEmpty(consumer, "consumer_tag"), + "active": boolOrDefault(consumer, "active", false), + "ackRequired": boolOrDefault(consumer, "ack_required", false), + } + if prefetch := integerOrNull(consumer, "prefetch_count"); prefetch != nil { + entry["prefetch"] = *prefetch + } + consumers = append(consumers, entry) + } + return consumers +} + +func (s *server) listNamespaces(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + vhosts, err := managementGet(connection, "/api/vhosts") + if err != nil { + return nil, err + } + array, ok := vhosts.([]any) + if !ok { + return nil, errors.New("Unexpected management API response for vhost listing") + } + namespaces := make([]jsonObject, 0, len(array)) + for _, value := range array { + vhost, ok := value.(map[string]any) + if ok { + namespaces = append(namespaces, jsonObject{"name": stringOrEmpty(jsonObject(vhost), "name")}) + } + } + return jsonObject{"namespaces": namespaces}, nil +} + +func (s *server) createNamespace(params jsonObject) (any, error) { + namespace, err := namespaceName(params) + if err != nil { + return nil, err + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + if _, err := managementSend(connection, http.MethodPut, "/api/vhosts/"+urlEncodeVhost(namespace), nil); err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) deleteNamespace(params jsonObject) (any, error) { + namespace, err := namespaceName(params) + if err != nil { + return nil, err + } + if err := assertNamespaceDeletable(namespace, ""); err != nil { + return nil, err + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + if err := assertNamespaceDeletable(namespace, stringOrDefault(connection, "virtual_host", "/")); err != nil { + return nil, err + } + if _, err := managementSend(connection, http.MethodDelete, "/api/vhosts/"+urlEncodeVhost(namespace), nil); err != nil { + return nil, err + } + return okResult(), nil +} + +func namespaceName(params jsonObject) (string, error) { + name := stringOrEmpty(params, "namespace") + if strings.TrimSpace(name) == "" { + return "", errors.New("namespace is required") + } + if strings.TrimSpace(name) == "*" { + return "", errors.New("namespace create/delete requires a specific virtual host (all-vhosts context)") + } + return name, nil +} + +func assertNamespaceDeletable(namespace, connectedVhost string) error { + if namespace == "/" { + return errors.New("The default virtual host '/' cannot be deleted") + } + if connectedVhost != "" && namespace == connectedVhost { + return fmt.Errorf("Cannot delete the virtual host '%s' while connected to it", namespace) + } + return nil +} + +func (s *server) listExchanges(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + allVhosts := allVhostsRequested(params) + exchanges, err := managementGetAll(connection, managementListPath(params, connection, "exchanges")) + if err != nil { + return nil, err + } + result := make([]jsonObject, 0, len(exchanges)) + for _, value := range exchanges { + exchange, ok := value.(map[string]any) + if !ok { + continue + } + info := exchangeInfoFromJSON(jsonObject(exchange)) + if allVhosts { + attachVhost(info, jsonObject(exchange)) + } + result = append(result, info) + } + sort.SliceStable(result, func(left, right int) bool { + return stringOrEmpty(result[left], "name") < stringOrEmpty(result[right], "name") + }) + return jsonObject{"exchanges": result}, nil +} + +func exchangeInfoFromJSON(exchange jsonObject) jsonObject { + exchangeType := stringOrEmpty(exchange, "type") + if exchangeType == "" { + exchangeType = "default" + } + return jsonObject{ + "name": stringOrEmpty(exchange, "name"), + "type": exchangeType, + "durable": boolOrDefault(exchange, "durable", false), + "autoDelete": boolOrDefault(exchange, "auto_delete", false), + "internal": boolOrDefault(exchange, "internal", false), + } +} + +func (s *server) createExchange(params jsonObject) (any, error) { + name, err := exchangeName(params) + if err != nil { + return nil, err + } + exchangeType, err := validateExchangeType(stringOrEmpty(params, "type")) + if err != nil { + return nil, err + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + vhost := effectiveVhost(params, connection) + body := jsonObject{ + "type": exchangeType, + "durable": boolOrDefault(params, "durable", true), + "auto_delete": boolOrDefault(params, "autoDelete", false), + } + if _, err := managementSend(connection, http.MethodPut, + "/api/exchanges/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(name), body); err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) deleteExchange(params jsonObject) (any, error) { + name := stringOrEmpty(params, "name") + if err := assertExchangeDeletable(name); err != nil { + return nil, err + } + if strings.TrimSpace(name) == "" { + return nil, errors.New("name is required") + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + vhost := effectiveVhost(params, connection) + if _, err := managementSend(connection, http.MethodDelete, + "/api/exchanges/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(name), nil); err != nil { + return nil, err + } + return okResult(), nil +} + +func exchangeName(params jsonObject) (string, error) { + name := stringOrEmpty(params, "name") + if strings.TrimSpace(name) == "" { + return "", errors.New("name is required") + } + return name, nil +} + +func validateExchangeType(exchangeType string) (string, error) { + if _, ok := exchangeTypes[exchangeType]; !ok { + return "", fmt.Errorf("Invalid exchange type '%s'. Supported types: direct, fanout, topic, headers", exchangeType) + } + return exchangeType, nil +} + +func assertExchangeDeletable(name string) error { + if name == "" { + return errors.New("The default exchange cannot be deleted") + } + if strings.HasPrefix(name, "amq.") { + return fmt.Errorf("The built-in exchange '%s' cannot be deleted", name) + } + return nil +} + +func (s *server) listBindings(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + allVhosts := allVhostsRequested(params) + bindings, err := managementGetAll(connection, managementListPath(params, connection, "bindings")) + if err != nil { + return nil, err + } + exchangeFilter := stringOrEmpty(params, "exchange") + queueFilter := stringOrEmpty(params, "queue") + result := make([]jsonObject, 0, len(bindings)) + for _, value := range bindings { + binding, ok := value.(map[string]any) + if !ok { + continue + } + info := bindingInfoFromJSON(jsonObject(binding)) + if exchangeFilter != "" && exchangeFilter != stringOrEmpty(info, "source") { + continue + } + if queueFilter != "" && !(queueFilter == stringOrEmpty(info, "destination") && stringOrEmpty(info, "destinationType") == "queue") { + continue + } + if allVhosts { + attachVhost(info, jsonObject(binding)) + } + result = append(result, info) + } + return jsonObject{"bindings": result}, nil +} + +func bindingInfoFromJSON(binding jsonObject) jsonObject { + info := jsonObject{ + "source": stringOrEmpty(binding, "source"), + "destination": stringOrEmpty(binding, "destination"), + "destinationType": stringOrEmpty(binding, "destination_type"), + "routingKey": stringOrEmpty(binding, "routing_key"), + } + if arguments := objectOrNil(binding, "arguments"); len(arguments) > 0 { + mapped := jsonObject{} + for key, value := range arguments { + if value == nil { + continue + } + if converted := argumentValue(value); converted != nil { + mapped[key] = converted + } else { + encoded, _ := json.Marshal(value) + mapped[key] = string(encoded) + } + } + if len(mapped) > 0 { + info["arguments"] = mapped + } + } + return info +} + +func (s *server) bind(params jsonObject) (any, error) { + if err := s.applyBinding(params, true); err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) unbind(params jsonObject) (any, error) { + if err := s.applyBinding(params, false); err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) applyBinding(params jsonObject, bind bool) error { + source, err := requireBindingName(params, "source") + if err != nil { + return err + } + destination, err := requireBindingName(params, "destination") + if err != nil { + return err + } + destinationType := stringOrDefault(params, "destinationType", stringOrEmpty(params, "destination_type")) + if destinationType != "queue" && destinationType != "exchange" { + return fmt.Errorf("destinationType must be 'queue' or 'exchange', got '%s'", destinationType) + } + routingKey := stringOrDefault(params, "routingKey", stringOrEmpty(params, "routing_key")) + arguments := bindingArguments(params) + channel, err := s.channelFor(params) + if err != nil { + return err + } + if destinationType == "queue" { + if bind { + return channel.QueueBind(destination, routingKey, source, false, arguments) + } + return channel.QueueUnbind(destination, routingKey, source, arguments) + } + if bind { + return channel.ExchangeBind(destination, routingKey, source, false, arguments) + } + return channel.ExchangeUnbind(destination, routingKey, source, false, arguments) +} + +func requireBindingName(params jsonObject, key string) (string, error) { + name := stringOrEmpty(params, key) + if strings.TrimSpace(name) == "" { + return "", fmt.Errorf("%s is required", key) + } + return name, nil +} + +func bindingArguments(params jsonObject) amqp.Table { + arguments := amqp.Table{} + if values := objectOrNil(params, "arguments"); values != nil { + for key, value := range values { + if converted := argumentValue(value); converted != nil { + arguments[key] = converted + } + } + } + return arguments +} + +func (s *server) listClientConnections(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + connections, err := managementGetAll(connection, "/api/connections") + if err != nil { + return nil, err + } + allVhosts := allVhostsRequested(params) + filter := vhostFilter(params, connection) + result := make([]jsonObject, 0, len(connections)) + for _, value := range connections { + entry, ok := value.(map[string]any) + if !ok { + continue + } + object := jsonObject(entry) + if filter != "" && filter != stringOrEmpty(object, "vhost") { + continue + } + info := clientConnectionInfoFromJSON(object) + if allVhosts { + attachVhost(info, object) + } + result = append(result, info) + } + sort.SliceStable(result, func(left, right int) bool { + return stringOrEmpty(result[left], "name") < stringOrEmpty(result[right], "name") + }) + return jsonObject{"connections": result}, nil +} + +func clientConnectionInfoFromJSON(connection jsonObject) jsonObject { + info := jsonObject{ + "name": stringOrEmpty(connection, "name"), + "user": stringOrEmpty(connection, "user"), + "peerHost": stringOrEmpty(connection, "peer_host"), + "peerPort": longOrDefault(connection, "peer_port", 0), + "state": stringOrEmpty(connection, "state"), + "channels": longOrDefault(connection, "channels", 0), + } + if rate := rateFromDetails(connection, "recv_oct_details"); rate != nil { + info["recvRate"] = *rate + } + if rate := rateFromDetails(connection, "send_oct_details"); rate != nil { + info["sendRate"] = *rate + } + if connectedAt := longOrNull(connection, "connected_at"); connectedAt != nil { + info["connectedAt"] = *connectedAt + } + return info +} + +func rateFromDetails(object jsonObject, key string) *float64 { + return floatOrNull(objectOrNil(object, key), "rate") +} + +func (s *server) listClientChannels(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + channels, err := managementGetAll(connection, "/api/channels") + if err != nil { + return nil, err + } + connectionFilter := stringOrEmpty(params, "connection") + allVhosts := allVhostsRequested(params) + filter := vhostFilter(params, connection) + result := make([]jsonObject, 0, len(channels)) + for _, value := range channels { + entry, ok := value.(map[string]any) + if !ok { + continue + } + object := jsonObject(entry) + if filter != "" && filter != stringOrEmpty(object, "vhost") { + continue + } + info := channelInfoFromJSON(object) + if allVhosts { + attachVhost(info, object) + } + if connectionFilter != "" && !channelMatchesConnection(info, connectionFilter) { + continue + } + result = append(result, info) + } + sort.SliceStable(result, func(left, right int) bool { + return stringOrEmpty(result[left], "name") < stringOrEmpty(result[right], "name") + }) + return jsonObject{"channels": result}, nil +} + +func channelInfoFromJSON(channel jsonObject) jsonObject { + info := jsonObject{ + "name": stringOrEmpty(channel, "name"), + "state": stringOrEmpty(channel, "state"), + } + if details := objectOrNil(channel, "connection_details"); details != nil { + if connectionName := stringOrEmpty(details, "name"); connectionName != "" { + info["connectionName"] = connectionName + } + } + if prefetch := integerOrNull(channel, "prefetch_count"); prefetch != nil { + info["prefetch"] = *prefetch + } + if unacked := longOrNull(channel, "messages_unacknowledged"); unacked != nil { + info["messagesUnacked"] = *unacked + } + if consumers := longOrNull(channel, "consumer_count"); consumers != nil { + info["consumerCount"] = *consumers + } + return info +} + +func channelMatchesConnection(channelInfo jsonObject, connectionName string) bool { + return stringOrEmpty(channelInfo, "connectionName") == connectionName || strings.HasPrefix(stringOrEmpty(channelInfo, "name"), connectionName) +} + +func (s *server) closeClientConnection(params jsonObject) (any, error) { + name := stringOrEmpty(params, "name") + if strings.TrimSpace(name) == "" { + return nil, errors.New("name is required") + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + if _, err := managementSend(connection, http.MethodDelete, "/api/connections/"+urlEncodeName(name), nil); err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) listUsers(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + users, err := managementGetAll(connection, "/api/users") + if err != nil { + return nil, err + } + result := make([]jsonObject, 0, len(users)) + for _, value := range users { + user, ok := value.(map[string]any) + if ok { + result = append(result, userInfoFromJSON(jsonObject(user))) + } + } + sort.SliceStable(result, func(left, right int) bool { + return stringOrEmpty(result[left], "name") < stringOrEmpty(result[right], "name") + }) + return jsonObject{"users": result}, nil +} + +func userInfoFromJSON(user jsonObject) jsonObject { + return jsonObject{ + "name": stringOrEmpty(user, "name"), + "tags": parseUserTags(stringOrEmpty(user, "tags")), + } +} + +func parseUserTags(tags string) []string { + result := make([]string, 0) + for _, tag := range strings.Split(tags, ",") { + if trimmed := strings.TrimSpace(tag); trimmed != "" { + result = append(result, trimmed) + } + } + return result +} + +func (s *server) createUser(params jsonObject) (any, error) { + name, err := userName(params) + if err != nil { + return nil, err + } + password := stringOrEmpty(params, "password") + if password == "" { + return nil, errors.New("password is required") + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + if err := assertNotConnectedUser("create or modify", name, stringOrDefault(connection, "username", "guest")); err != nil { + return nil, err + } + body := jsonObject{"password": password, "tags": userTagsParam(params)} + if _, err := managementSend(connection, http.MethodPut, "/api/users/"+urlEncodePathSegment(name), body); err != nil { + return nil, err + } + return okResult(), nil +} + +func userTagsParam(params jsonObject) string { + value, exists := params["tags"] + if !exists || value == nil { + return "" + } + if tags, ok := value.([]any); ok { + parts := make([]string, 0, len(tags)) + for _, tag := range tags { + trimmed := strings.TrimSpace(fmt.Sprint(tag)) + if trimmed != "" { + parts = append(parts, trimmed) + } + } + return strings.Join(parts, ",") + } + return fmt.Sprint(value) +} + +func (s *server) deleteUser(params jsonObject) (any, error) { + name, err := userName(params) + if err != nil { + return nil, err + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + if err := assertNotConnectedUser("delete", name, stringOrDefault(connection, "username", "guest")); err != nil { + return nil, err + } + if _, err := managementSend(connection, http.MethodDelete, "/api/users/"+urlEncodePathSegment(name), nil); err != nil { + return nil, err + } + return okResult(), nil +} + +func assertNotConnectedUser(action, name, connectedUser string) error { + if name == connectedUser { + return fmt.Errorf("Cannot %s user '%s' while connected as that user", action, name) + } + return nil +} + +func (s *server) listPermissions(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + permissions, err := managementGet(connection, "/api/permissions") + if err != nil { + return nil, err + } + array, ok := permissions.([]any) + if !ok { + return nil, errors.New("Unexpected management API response for permission listing") + } + vhostFilter := stringOrEmpty(params, "virtual_host") + if allVhostsRequested(params) { + vhostFilter = "" + } + userFilter := stringOrEmpty(params, "user") + result := make([]jsonObject, 0, len(array)) + for _, value := range array { + permission, ok := value.(map[string]any) + if !ok { + continue + } + info := permissionInfoFromJSON(jsonObject(permission)) + if vhostFilter != "" && vhostFilter != stringOrEmpty(info, "vhost") { + continue + } + if userFilter != "" && userFilter != stringOrEmpty(info, "user") { + continue + } + result = append(result, info) + } + sort.SliceStable(result, func(left, right int) bool { + leftUser := stringOrEmpty(result[left], "user") + rightUser := stringOrEmpty(result[right], "user") + if leftUser != rightUser { + return leftUser < rightUser + } + return stringOrEmpty(result[left], "vhost") < stringOrEmpty(result[right], "vhost") + }) + return jsonObject{"permissions": result}, nil +} + +func permissionInfoFromJSON(permission jsonObject) jsonObject { + return jsonObject{ + "user": stringOrEmpty(permission, "user"), + "vhost": stringOrEmpty(permission, "vhost"), + "configure": stringOrEmpty(permission, "configure"), + "write": stringOrEmpty(permission, "write"), + "read": stringOrEmpty(permission, "read"), + } +} + +func (s *server) grantPermission(params jsonObject) (any, error) { + user, err := userName(params) + if err != nil { + return nil, err + } + vhost, err := permissionVhost(params) + if err != nil { + return nil, err + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + body := jsonObject{ + "configure": permissionPattern(params, "configure"), + "write": permissionPattern(params, "write"), + "read": permissionPattern(params, "read"), + } + if _, err := managementSend(connection, http.MethodPut, + "/api/permissions/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(user), body); err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) revokePermission(params jsonObject) (any, error) { + user, err := userName(params) + if err != nil { + return nil, err + } + vhost, err := permissionVhost(params) + if err != nil { + return nil, err + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + if _, err := managementSend(connection, http.MethodDelete, + "/api/permissions/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(user), nil); err != nil { + return nil, err + } + return okResult(), nil +} + +func permissionPattern(params jsonObject, key string) string { + pattern := stringOrEmpty(params, key) + if strings.TrimSpace(pattern) == "" { + return defaultPermissionPattern + } + return pattern +} + +func permissionVhost(params jsonObject) (string, error) { + vhost := stringOrEmpty(params, "virtual_host") + if strings.TrimSpace(vhost) == "" { + return "", errors.New("virtual_host is required") + } + if vhost == "*" { + return "", errors.New("all_vhosts is only supported for list operations") + } + return vhost, nil +} + +func userName(params jsonObject) (string, error) { + name := stringOrEmpty(params, "name") + if strings.TrimSpace(name) == "" { + name = stringOrEmpty(params, "user") + } + if strings.TrimSpace(name) == "" { + return "", errors.New("user name is required") + } + return name, nil +} + +func (s *server) listPolicies(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + allVhosts := allVhostsRequested(params) || stringOrEmpty(params, "virtual_host") == "*" + path := "/api/policies/" + urlEncodeVhost(effectiveVhost(params, connection)) + if allVhosts { + path = "/api/policies" + } + policies, err := managementGetAll(connection, path) + if err != nil { + return nil, err + } + result := make([]jsonObject, 0, len(policies)) + for _, value := range policies { + policy, ok := value.(map[string]any) + if ok { + result = append(result, policyInfoFromJSON(jsonObject(policy))) + } + } + sort.SliceStable(result, func(left, right int) bool { + leftVhost := stringOrEmpty(result[left], "vhost") + rightVhost := stringOrEmpty(result[right], "vhost") + if leftVhost != rightVhost { + return leftVhost < rightVhost + } + return stringOrEmpty(result[left], "name") < stringOrEmpty(result[right], "name") + }) + return jsonObject{"policies": result}, nil +} + +func policyInfoFromJSON(policy jsonObject) jsonObject { + definition := jsonObject{} + if rawDefinition := objectOrNil(policy, "definition"); rawDefinition != nil { + for key, value := range rawDefinition { + if value == nil { + continue + } + if converted := argumentValue(value); converted != nil { + definition[key] = converted + } else { + encoded, _ := json.Marshal(value) + definition[key] = string(encoded) + } + } + } + return jsonObject{ + "name": stringOrEmpty(policy, "name"), + "vhost": stringOrEmpty(policy, "vhost"), + "pattern": stringOrEmpty(policy, "pattern"), + "applyTo": stringOrEmpty(policy, "apply-to"), + "priority": longOrDefault(policy, "priority", 0), + "definition": definition, + } +} + +func (s *server) setPolicy(params jsonObject) (any, error) { + vhost, err := permissionVhost(params) + if err != nil { + return nil, err + } + name, err := policyName(params) + if err != nil { + return nil, err + } + pattern := stringOrEmpty(params, "pattern") + if strings.TrimSpace(pattern) == "" { + return nil, errors.New("pattern is required") + } + definition := objectOrNil(params, "definition") + if definition == nil { + return nil, errors.New("definition is required") + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + body := jsonObject{ + "pattern": pattern, + "apply-to": stringOrDefault(params, "applyTo", "queues"), + "priority": intOrDefault(params, "priority", 0), + "definition": definition, + } + if _, err := managementSend(connection, http.MethodPut, + "/api/policies/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(name), body); err != nil { + return nil, err + } + return okResult(), nil +} + +func (s *server) deletePolicy(params jsonObject) (any, error) { + vhost, err := permissionVhost(params) + if err != nil { + return nil, err + } + name, err := policyName(params) + if err != nil { + return nil, err + } + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + if _, err := managementSend(connection, http.MethodDelete, + "/api/policies/"+urlEncodeVhost(vhost)+"/"+urlEncodePathSegment(name), nil); err != nil { + return nil, err + } + return okResult(), nil +} + +func policyName(params jsonObject) (string, error) { + name := stringOrEmpty(params, "name") + if strings.TrimSpace(name) == "" { + return "", errors.New("name is required") + } + return name, nil +} + +func (s *server) peekMessages(params jsonObject) (any, error) { + queue, err := queueName(params) + if err != nil { + return nil, err + } + offset := normalizePeekOffset(longOrDefault(params, "offset", 0)) + count := normalizePeekCount(intOrDefault(params, "count", 10)) + if offset >= maxPeekMessages { + return jsonObject{"messages": []jsonObject{}}, nil + } + totalToFetch := offset + int64(count) + if totalToFetch > maxPeekMessages { + totalToFetch = maxPeekMessages + } + var channel *amqp.Channel + var ownedConnection *amqp.Connection + if s.cachedConnection != nil { + channel, err = s.channelFor(params) + } else { + config := deepCopyObject(connectionObject(params)) + if vhost := stringOrNull(params, "virtual_host"); vhost != nil && strings.TrimSpace(*vhost) != "" { + config["virtual_host"] = *vhost + } + ownedConnection, err = openConnection(config) + if err == nil { + channel, err = ownedConnection.Channel() + } + } + if err != nil { + closeConnection(ownedConnection) + return nil, err + } + if ownedConnection != nil { + defer closeConnection(ownedConnection) + defer closeChannel(channel) + } + fetched := make([]amqp.Delivery, 0, totalToFetch) + var lastDeliveryTag uint64 + for index := int64(0); index < totalToFetch; index++ { + delivery, ok, getError := channel.Get(queue, false) + if getError != nil { + return nil, getError + } + if !ok { + break + } + fetched = append(fetched, delivery) + lastDeliveryTag = delivery.DeliveryTag + } + if lastDeliveryTag > 0 { + if err := channel.Nack(lastDeliveryTag, true, true); err != nil { + return nil, err + } + } + messageCapacity := peekMessageCapacity(len(fetched), offset, count) + messages := make([]jsonObject, 0, messageCapacity) + for index := offset; index < int64(len(fetched)) && len(messages) < count; index++ { + messages = append(messages, peekedMessageFromDelivery(queue, index, fetched[index])) + } + return jsonObject{"messages": messages}, nil +} + +func peekMessageCapacity(fetched int, offset int64, count int) int { + available := fetched - int(offset) + if available < 0 { + return 0 + } + if available > count { + return count + } + return available +} + +func normalizePeekOffset(requested int64) int64 { + if requested < 0 { + return 0 + } + return requested +} + +func normalizePeekCount(requested int) int { + if requested < 1 { + return 1 + } + return requested +} + +func resolveRoutingKey(params jsonObject, queue string) string { + routingKey := stringOrDefault(params, "routing_key", "") + if routingKey == "" { + routingKey = stringOrDefault(params, "routingKey", "") + } + if routingKey == "" { + routingKey = stringOrDefault(params, "key", "") + } + if strings.TrimSpace(routingKey) == "" { + return queue + } + return routingKey +} + +func peekedMessageFromDelivery(queue string, index int64, delivery amqp.Delivery) jsonObject { + message := jsonObject{ + "topic": queue, + "offset": index, + "exchange": delivery.Exchange, + "routingKey": delivery.RoutingKey, + "redelivered": delivery.Redelivered, + "deliveryTag": delivery.DeliveryTag, + "timestamp": int64(0), + "headers": stringHeaders(delivery.Headers), + "payloadBase64": base64.StdEncoding.EncodeToString(delivery.Body), + } + if delivery.MessageId != "" { + message["messageId"] = delivery.MessageId + } + if !delivery.Timestamp.IsZero() { + message["timestamp"] = delivery.Timestamp.UnixMilli() + } + if utf8.Valid(delivery.Body) { + message["payloadText"] = string(delivery.Body) + } + return message +} + +func stringHeaders(headers amqp.Table) map[string]string { + result := make(map[string]string, len(headers)) + for key, value := range headers { + switch typed := value.(type) { + case []byte: + result[key] = string(typed) + default: + result[key] = fmt.Sprint(typed) + } + } + return result +} + +func (s *server) sendMessage(params jsonObject) (any, error) { + channel, err := s.channelFor(params) + if err != nil { + return nil, err + } + queue, err := queueName(params) + if err != nil { + return nil, err + } + exchange := stringOrDefault(params, "exchange", "") + routingKey := resolveRoutingKey(params, queue) + payload := stringOrEmpty(params, "payloadBase64") + body := []byte{} + if payload != "" { + body, err = base64.StdEncoding.DecodeString(payload) + if err != nil { + return nil, err + } + } + publishing := amqp.Publishing{Body: body} + if headers := objectOrNil(params, "headers"); headers != nil { + publishing.Headers = amqp.Table{} + for key, value := range headers { + if converted := argumentValue(value); converted != nil { + publishing.Headers[key] = converted + } + } + } + if err := channel.PublishWithContext(context.Background(), exchange, routingKey, false, false, publishing); err != nil { + return nil, err + } + return jsonObject{"ok": true, "exchange": exchange, "routingKey": routingKey}, nil +} + +func (s *server) describeCluster(params jsonObject) (any, error) { + connection, err := s.requireConnection() + if err != nil { + return nil, err + } + config := s.currentConnectionConfig(params) + nodes := make([]jsonObject, 0) + if config != nil { + addresses, resolveError := resolveAddresses(config) + if resolveError != nil { + return nil, resolveError + } + for _, endpoint := range addresses { + nodes = append(nodes, jsonObject{"name": endpoint.Host, "port": endpoint.Port}) + } + } + return jsonObject{ + "clusterName": serverString(connection.Properties, "cluster_name"), + "product": serverString(connection.Properties, "product"), + "version": serverString(connection.Properties, "version"), + "platform": serverString(connection.Properties, "platform"), + "nodes": nodes, + "nodeCount": len(nodes), + }, nil +} + +func (s *server) getOverview(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + overview, err := managementGet(connection, "/api/overview") + if err != nil { + return nil, err + } + object, ok := overview.(map[string]any) + if !ok { + return nil, errors.New("Unexpected management API response for cluster overview") + } + return overviewInfoFromJSON(jsonObject(object)), nil +} + +func overviewInfoFromJSON(overview jsonObject) jsonObject { + info := jsonObject{} + putIfPresent(info, "messagesReady", nestedLongOrNull(overview, "queue_totals", "messages_ready")) + putIfPresent(info, "messagesUnacked", nestedLongOrNull(overview, "queue_totals", "messages_unacknowledged")) + if stats := objectOrNil(overview, "message_stats"); stats != nil { + putIfPresent(info, "publishRate", rateFromDetails(stats, "publish_details")) + putIfPresent(info, "deliverRate", rateFromDetails(stats, "deliver_get_details")) + putIfPresent(info, "ackRate", rateFromDetails(stats, "ack_details")) + } + putIfPresent(info, "totalQueues", nestedLongOrNull(overview, "object_totals", "queues")) + putIfPresent(info, "totalExchanges", nestedLongOrNull(overview, "object_totals", "exchanges")) + putIfPresent(info, "totalConnections", nestedLongOrNull(overview, "object_totals", "connections")) + putIfPresent(info, "totalChannels", nestedLongOrNull(overview, "object_totals", "channels")) + putIfPresent(info, "totalConsumers", nestedLongOrNull(overview, "object_totals", "consumers")) + return info +} + +func (s *server) listNodes(params jsonObject) (any, error) { + connection, err := s.requireConnectionConfig(params) + if err != nil { + return nil, err + } + nodes, err := managementGet(connection, "/api/nodes") + if err != nil { + return nil, err + } + array, ok := nodes.([]any) + if !ok { + return nil, errors.New("Unexpected management API response for node listing") + } + result := make([]jsonObject, 0, len(array)) + for _, value := range array { + node, ok := value.(map[string]any) + if ok { + result = append(result, nodeInfoFromJSON(jsonObject(node))) + } + } + sort.SliceStable(result, func(left, right int) bool { + return stringOrEmpty(result[left], "name") < stringOrEmpty(result[right], "name") + }) + return jsonObject{"nodes": result}, nil +} + +func nodeInfoFromJSON(node jsonObject) jsonObject { + info := jsonObject{ + "name": stringOrEmpty(node, "name"), + "running": boolOrDefault(node, "running", false), + } + putIfPresent(info, "memUsed", longOrNull(node, "mem_used")) + putIfPresent(info, "memLimit", longOrNull(node, "mem_limit")) + putIfPresent(info, "diskFree", longOrNull(node, "disk_free")) + putIfPresent(info, "fdUsed", longOrNull(node, "fd_used")) + putIfPresent(info, "fdTotal", longOrNull(node, "fd_total")) + putIfPresent(info, "socketsUsed", longOrNull(node, "sockets_used")) + putIfPresent(info, "socketsTotal", longOrNull(node, "sockets_total")) + putIfPresent(info, "uptimeMs", longOrNull(node, "uptime")) + return info +} + +func nestedLongOrNull(object jsonObject, block, key string) *int64 { + return longOrNull(objectOrNil(object, block), key) +} + +func putIfPresent(info jsonObject, key string, value any) { + switch typed := value.(type) { + case *int64: + if typed != nil { + info[key] = *typed + } + case *float64: + if typed != nil { + info[key] = *typed + } + case nil: + default: + info[key] = typed + } +} diff --git a/agents/drivers/rabbitmq/src/main/java/com/dbx/agent/rabbitmq/RabbitMqAgent.java b/agents/drivers/rabbitmq/src/main/java/com/dbx/agent/rabbitmq/RabbitMqAgent.java deleted file mode 100644 index 8f71465e5..000000000 --- a/agents/drivers/rabbitmq/src/main/java/com/dbx/agent/rabbitmq/RabbitMqAgent.java +++ /dev/null @@ -1,2236 +0,0 @@ -package com.dbx.agent.rabbitmq; - -import com.google.gson.*; -import com.rabbitmq.client.AMQP; -import com.rabbitmq.client.Address; -import com.rabbitmq.client.Channel; -import com.rabbitmq.client.Connection; -import com.rabbitmq.client.ConnectionFactory; -import com.rabbitmq.client.GetResponse; -import com.rabbitmq.client.ShutdownSignalException; - -import javax.net.ssl.HttpsURLConnection; -import javax.net.ssl.SSLContext; -import javax.net.ssl.TrustManager; -import javax.net.ssl.X509TrustManager; -import java.io.BufferedReader; -import java.io.IOException; -import java.io.InputStream; -import java.io.InputStreamReader; -import java.io.OutputStream; -import java.io.PrintStream; -import java.net.HttpURLConnection; -import java.net.URI; -import java.net.URL; -import java.net.URLEncoder; -import java.nio.charset.StandardCharsets; -import java.security.SecureRandom; -import java.security.cert.X509Certificate; -import java.util.*; -import java.util.regex.Matcher; -import java.util.regex.Pattern; - -/** - * RabbitMQ admin agent for DBX. Communicates with the Rust bridge via JSON-RPC - * over stdin/stdout. Uses the RabbitMQ AMQP Java client for queue operations - * and the HTTP management API (when available) for queue listing. - */ -public final class RabbitMqAgent { - - private static final Gson GSON = new GsonBuilder().serializeNulls().create(); - private static final int DEFAULT_PORT = 5672; - private static final int DEFAULT_REQUEST_TIMEOUT_MS = 30_000; - private static final int DEFAULT_MANAGEMENT_PORT = 15672; - private static final int DEFAULT_MANAGEMENT_TLS_PORT = 15671; - private static final int MAX_PEEK_MESSAGES = 10_000; - - private static final List CAPABILITIES = Collections.unmodifiableList(Arrays.asList( - "mq_connect", "mq_test_connection", "mq_topics", - "mq_messages", "mq_config", "mq_monitoring", "mq_exchanges", - "mq_client_connections", "mq_user_permissions", "mq_policies" - )); - - private static Connection connection; - private static Channel channel; - private static JsonObject cachedConnection; - // Lazily created AMQP clients for virtual hosts other than the connection's - // default vhost; AMQP connections are scoped to a single vhost, so each - // extra vhost needs its own connection/channel pair. - private static final Map vhostClients = new HashMap<>(); - private static volatile boolean shutdownRequested; - - private RabbitMqAgent() {} - - // ----------------------------------------------------------------------- - // Entry point - // ----------------------------------------------------------------------- - - public static void main(String[] args) throws Exception { - System.setProperty("org.slf4j.simpleLogger.logFile", "System.err"); - // The JSON-RPC pipe with the Rust bridge is UTF-8; relying on the - // platform default charset mangles non-ASCII payloads on Windows. - PrintStream out = new PrintStream(System.out, true, StandardCharsets.UTF_8); - out.println("{\"ready\":true}"); - out.flush(); - - BufferedReader reader = new BufferedReader(new InputStreamReader(System.in, StandardCharsets.UTF_8)); - while (true) { - String line = reader.readLine(); - if (line == null) break; - String response = handleRequest(line); - out.println(response); - out.flush(); - if (shutdownRequested) { - System.exit(0); - } - } - } - - // ----------------------------------------------------------------------- - // JSON-RPC dispatch - // ----------------------------------------------------------------------- - - static String handleRequest(String line) { - JsonObject response = new JsonObject(); - response.addProperty("jsonrpc", "2.0"); - try { - // Parse and extract inside the try: a malformed request must yield - // a JSON-RPC error (with a null id), never kill the agent process. - JsonObject req = JsonParser.parseString(line).getAsJsonObject(); - JsonElement id = req.get("id"); - response.add("id", id != null ? id : JsonNull.INSTANCE); - String method = req.get("method").getAsString(); - JsonObject params = req.has("params") && req.get("params").isJsonObject() - ? req.getAsJsonObject("params") : new JsonObject(); - - Object result = dispatch(method, params); - response.add("result", GSON.toJsonTree(result)); - } catch (Exception e) { - if (!response.has("id")) { - response.add("id", JsonNull.INSTANCE); - } - JsonObject error = new JsonObject(); - error.addProperty("code", -1); - error.addProperty("message", normalizeErrorMessage(e)); - response.add("error", error); - } - return GSON.toJson(response); - } - - /** - * Operations that act on a single vhost-scoped resource. The {@code all_vhosts} - * sentinel only makes sense for cluster-wide listings; for these methods it - * must fail fast instead of silently falling back to the default vhost. - */ - private static final Set ALL_VHOSTS_UNSUPPORTED_METHODS = Set.of( - "mq_create_topic", "mq_delete_topic", "mq_purge_queue", "mq_send_message", - "mq_bind", "mq_unbind", "mq_create_exchange", "mq_delete_exchange", - "mq_peek_messages", "mq_get_topic_stats", "mq_list_consumers", "mq_close_connection", - "mq_grant_permission", "mq_revoke_permission", "mq_set_policy", "mq_delete_policy"); - - private static Object dispatch(String method, JsonObject params) throws Exception { - if (ALL_VHOSTS_UNSUPPORTED_METHODS.contains(method) && allVhostsRequested(params)) { - throw new IllegalArgumentException("all_vhosts is only supported for list operations"); - } - return switch (method) { - case "handshake" -> handshakeResult(); - case "connect" -> connect(params); - case "test_connection" -> testConnection(params); - case "disconnect" -> { closeClients(); yield Collections.singletonMap("ok", true); } - case "shutdown" -> { closeClients(); shutdownRequested = true; yield Collections.singletonMap("ok", true); } - // Topic (queue) management - case "mq_list_topics" -> listTopics(params); - case "mq_create_topic" -> createTopic(params); - case "mq_delete_topic" -> deleteTopic(params); - case "mq_get_topic_stats" -> getTopicStats(params); - case "mq_get_topic_config" -> getTopicConfig(params); - case "mq_alter_topic_config" -> alterTopicConfig(params); - case "mq_purge_queue" -> purgeQueue(params); - case "mq_list_consumers" -> listConsumers(params); - // Namespaces (virtual hosts) - case "mq_list_namespaces" -> listNamespaces(params); - case "mq_create_namespace" -> createNamespace(params); - case "mq_delete_namespace" -> deleteNamespace(params); - // Exchanges & bindings - case "mq_list_exchanges" -> listExchanges(params); - case "mq_create_exchange" -> createExchange(params); - case "mq_delete_exchange" -> deleteExchange(params); - case "mq_list_bindings" -> listBindings(params); - case "mq_bind" -> bind(params); - case "mq_unbind" -> unbind(params); - // Client connections & channels - case "mq_list_connections" -> listClientConnections(params); - case "mq_list_channels" -> listClientChannels(params); - case "mq_close_connection" -> closeClientConnection(params); - // Users & permissions - case "mq_list_users" -> listUsers(params); - case "mq_create_user" -> createUser(params); - case "mq_delete_user" -> deleteUser(params); - case "mq_list_permissions" -> listPermissions(params); - case "mq_grant_permission" -> grantPermission(params); - case "mq_revoke_permission" -> revokePermission(params); - // Policies - case "mq_list_policies" -> listPolicies(params); - case "mq_set_policy" -> setPolicy(params); - case "mq_delete_policy" -> deletePolicy(params); - // Messages - case "mq_peek_messages" -> peekMessages(params); - case "mq_send_message" -> sendMessage(params); - // Cluster / monitoring - case "mq_describe_cluster" -> describeCluster(params); - case "mq_overview" -> getOverview(params); - case "mq_list_nodes" -> listNodes(params); - default -> throw new IllegalArgumentException("Unknown method: " + method); - }; - } - - // ----------------------------------------------------------------------- - // Lifecycle - // ----------------------------------------------------------------------- - - private static Object handshakeResult() { - return new HandshakeResult(1, 1, CAPABILITIES); - } - - private static Object connect(JsonObject params) throws Exception { - JsonObject conn = connectionObject(params); - Connection nextConnection = null; - Channel nextChannel = null; - try { - nextConnection = openConnection(conn); - nextChannel = nextConnection.createChannel(); - closeClients(); - connection = nextConnection; - channel = nextChannel; - cachedConnection = conn.deepCopy(); - return Collections.singletonMap("ok", true); - } catch (Exception e) { - closeQuietly(nextChannel); - closeQuietly(nextConnection); - throw e; - } - } - - private static Object testConnection(JsonObject params) throws Exception { - JsonObject conn = connectionObject(params); - Connection probe = null; - try { - probe = openConnection(conn); - Map serverProps = probe.getServerProperties(); - - Map result = new LinkedHashMap<>(); - result.put("ok", true); - result.put("product", serverString(serverProps, "product")); - result.put("version", serverString(serverProps, "version")); - result.put("serverVersion", serverString(serverProps, "version")); - result.put("clusterName", serverString(serverProps, "cluster_name")); - result.put("platform", serverString(serverProps, "platform")); - return result; - } finally { - closeQuietly(probe); - } - } - - private static void closeClients() { - for (VhostClient client : vhostClients.values()) { - client.closeQuietly(); - } - vhostClients.clear(); - closeQuietly(channel); - channel = null; - closeQuietly(connection); - connection = null; - cachedConnection = null; - } - - private static void closeQuietly(Channel ch) { - if (ch != null) { - try { - ch.close(); - } catch (Exception ignored) {} - } - } - - private static void closeQuietly(Connection conn) { - if (conn != null) { - try { - conn.close(); - } catch (Exception ignored) {} - } - } - - // ----------------------------------------------------------------------- - // Client builders - // ----------------------------------------------------------------------- - - private static Connection openConnection(JsonObject conn) throws Exception { - ConnectionFactory factory = buildConnectionFactory(conn); - List
addresses = resolveAddresses(conn); - return factory.newConnection(addresses); - } - - static ConnectionFactory buildConnectionFactory(JsonObject conn) throws Exception { - ConnectionFactory factory = new ConnectionFactory(); - factory.setUsername(credentialOrGuest(conn, "username")); - factory.setPassword(credentialOrGuest(conn, "password")); - factory.setVirtualHost(stringOrDefault(conn, "virtual_host", "/")); - factory.setConnectionTimeout(intOrDefault(conn, "request_timeout_ms", DEFAULT_REQUEST_TIMEOUT_MS)); - applyTlsSettings(conn, factory); - applyExtraProperties(conn, factory); - return factory; - } - - static void applyTlsSettings(JsonObject conn, ConnectionFactory factory) throws Exception { - JsonObject tls = conn.has("tls") && conn.get("tls").isJsonObject() - ? conn.getAsJsonObject("tls") : null; - boolean tlsEnabled = tls != null - || boolOrDefault(conn, "tls_skip_verify", false) - || boolProperty(conn, "ssl") - || boolProperty(conn, "tls"); - if (!tlsEnabled) { - return; - } - boolean skipVerify = tlsSkipVerify(conn); - if (skipVerify) { - factory.useSslProtocol(trustAllSslContext()); - } else { - factory.useSslProtocol(); - factory.enableHostnameVerification(); - } - } - - static void applyExtraProperties(JsonObject conn, ConnectionFactory factory) { - JsonObject properties = conn.has("properties") && conn.get("properties").isJsonObject() - ? conn.getAsJsonObject("properties") : null; - if (properties == null) { - return; - } - Integer heartbeat = integerProperty(properties, "requested_heartbeat"); - if (heartbeat != null) { - factory.setRequestedHeartbeat(heartbeat); - } - Integer connectionTimeout = integerProperty(properties, "connection_timeout_ms"); - if (connectionTimeout != null) { - factory.setConnectionTimeout(connectionTimeout); - } - Integer handshakeTimeout = integerProperty(properties, "handshake_timeout_ms"); - if (handshakeTimeout != null) { - factory.setHandshakeTimeout(handshakeTimeout); - } - Boolean automaticRecovery = booleanProperty(properties, "automatic_recovery"); - if (automaticRecovery != null) { - factory.setAutomaticRecoveryEnabled(automaticRecovery); - } - Boolean topologyRecovery = booleanProperty(properties, "topology_recovery"); - if (topologyRecovery != null) { - factory.setTopologyRecoveryEnabled(topologyRecovery); - } - } - - /** - * Parse the {@code addresses} connection parameter: a comma-separated list of - * {@code host[:port]} entries. Bare hosts fall back to {@code defaultPort} - * (the {@code port} connection parameter, defaulting to 5672). - */ - static List
resolveAddresses(JsonObject conn) { - String addresses = stringOrEmpty(conn, "addresses"); - if (addresses.isBlank()) { - addresses = stringOrEmpty(conn, "host"); - } - if (addresses.isBlank()) { - throw new IllegalArgumentException("addresses is required"); - } - return parseAddresses(addresses, intOrDefault(conn, "port", DEFAULT_PORT)); - } - - static List
parseAddresses(String addresses, int defaultPort) { - List
result = new ArrayList<>(); - for (String part : addresses.split(",")) { - String trimmed = part.trim(); - if (trimmed.isEmpty()) { - continue; - } - int colon = trimmed.lastIndexOf(':'); - if (colon > 0 && colon < trimmed.length() - 1) { - result.add(new Address(trimmed.substring(0, colon), Integer.parseInt(trimmed.substring(colon + 1)))); - } else { - result.add(new Address(trimmed, defaultPort)); - } - } - if (result.isEmpty()) { - throw new IllegalArgumentException("addresses is required"); - } - return result; - } - - /** Whether the connection config asks to skip TLS certificate verification. */ - static boolean tlsSkipVerify(JsonObject conn) { - JsonObject tls = conn.has("tls") && conn.get("tls").isJsonObject() - ? conn.getAsJsonObject("tls") : null; - return boolOrDefault(conn, "tls_skip_verify", false) - || (tls != null && boolOrDefault(tls, "skip_verify", false)); - } - - private static SSLContext trustAllSslContext() throws Exception { - TrustManager[] trustAll = new TrustManager[] { - new X509TrustManager() { - @Override - public void checkClientTrusted(X509Certificate[] chain, String authType) {} - - @Override - public void checkServerTrusted(X509Certificate[] chain, String authType) {} - - @Override - public X509Certificate[] getAcceptedIssuers() { - return new X509Certificate[0]; - } - } - }; - SSLContext context = SSLContext.getInstance("TLS"); - context.init(null, trustAll, new SecureRandom()); - return context; - } - - // ----------------------------------------------------------------------- - // Topic (queue) management - // ----------------------------------------------------------------------- - - private static Object listTopics(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - boolean allVhosts = allVhostsRequested(params); - JsonArray queues = managementGetAll(conn, managementListPath(params, conn, "queues")); - - List> topics = new ArrayList<>(); - for (JsonElement element : queues) { - JsonObject queue = element.getAsJsonObject(); - Map topic = new LinkedHashMap<>(); - topic.put("name", stringOrEmpty(queue, "name")); - topic.put("durable", boolOrDefault(queue, "durable", false)); - topic.put("autoDelete", boolOrDefault(queue, "auto_delete", false)); - topic.put("state", stringOrEmpty(queue, "state")); - topic.put("messages", longOrDefault(queue, "messages", 0)); - topic.put("consumers", longOrDefault(queue, "consumers", 0)); - if (allVhosts) { - attachVhost(topic, queue); - } - topics.add(topic); - } - topics.sort(Comparator.comparing(m -> (String) m.get("name"))); - return Collections.singletonMap("topics", topics); - } - - private static Object createTopic(JsonObject params) throws Exception { - Channel ch = channelFor(params); - String name = queueName(params); - boolean durable = boolOrDefault(params, "durable", true); - - Map arguments = new HashMap<>(); - JsonObject configs = params.has("configs") && params.get("configs").isJsonObject() - ? params.getAsJsonObject("configs") : null; - if (configs != null) { - for (Map.Entry entry : configs.entrySet()) { - Object value = argumentValue(entry.getValue()); - if (value != null) { - arguments.put(entry.getKey(), value); - } - } - } - - ch.queueDeclare(name, durable, false, false, arguments); - return Collections.singletonMap("ok", true); - } - - private static Object deleteTopic(JsonObject params) throws Exception { - Channel ch = channelFor(params); - ch.queueDelete(queueName(params)); - return Collections.singletonMap("ok", true); - } - - private static Object getTopicStats(JsonObject params) throws Exception { - String name = queueName(params); - - // Prefer the management API: it is read-only and works for exclusive - // queues, whereas a passive declare on an exclusive queue owned by - // another connection fails with 405 RESOURCE_LOCKED and the broker - // force-closes the channel. - JsonObject conn = currentConnectionConfig(params); - if (conn != null) { - String vhost = effectiveVhost(params, conn); - try { - JsonElement queue = managementGet(conn, - "/api/queues/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(name)); - if (queue.isJsonObject()) { - JsonObject info = queue.getAsJsonObject(); - long messages = longOrDefault(info, "messages", 0); - Map result = new LinkedHashMap<>(); - result.put("name", name); - result.put("messageCount", messages); - result.put("consumerCount", longOrDefault(info, "consumers", 0)); - result.put("totalMessages", messages); - return result; - } - } catch (Exception managementError) { - System.err.println("Management API unavailable for queue stats, " - + "falling back to passive declare: " + managementError.getMessage()); - } - } - - Channel ch = channelFor(params); - AMQP.Queue.DeclareOk declared = ch.queueDeclarePassive(name); - - Map result = new LinkedHashMap<>(); - result.put("name", name); - result.put("messageCount", declared.getMessageCount()); - result.put("consumerCount", declared.getConsumerCount()); - result.put("totalMessages", declared.getMessageCount()); - return result; - } - - private static Object getTopicConfig(JsonObject params) throws Exception { - Channel ch = channelFor(params); - String name = queueName(params); - - Map configs = new LinkedHashMap<>(); - JsonObject conn = currentConnectionConfig(params); - if (conn != null) { - String vhost = effectiveVhost(params, conn); - try { - JsonElement queue = managementGet(conn, - "/api/queues/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(name)); - if (queue.isJsonObject()) { - JsonObject info = queue.getAsJsonObject(); - configs.put("durable", boolOrDefault(info, "durable", false)); - configs.put("auto_delete", boolOrDefault(info, "auto_delete", false)); - configs.put("exclusive", boolOrDefault(info, "exclusive", false)); - if (info.has("arguments") && info.get("arguments").isJsonObject()) { - for (Map.Entry entry : info.getAsJsonObject("arguments").entrySet()) { - configs.put(entry.getKey(), entry.getValue().isJsonNull() - ? null : entry.getValue().getAsString()); - } - } - } - } catch (Exception managementError) { - System.err.println("Management API unavailable for queue config: " + managementError.getMessage()); - } - } - - // Fall back to a passive declare so the call still verifies the queue exists. - if (configs.isEmpty()) { - ch.queueDeclarePassive(name); - } - return Collections.singletonMap("configs", configs); - } - - private static Object alterTopicConfig(JsonObject params) { - throw new UnsupportedOperationException( - "RabbitMQ queue arguments are immutable after declaration; delete and re-declare the queue to change them"); - } - - private static Object purgeQueue(JsonObject params) throws Exception { - String name = queueName(params); - Channel ch = channelFor(params); - AMQP.Queue.PurgeOk purged = ch.queuePurge(name); - - Map result = new LinkedHashMap<>(); - result.put("ok", true); - result.put("purged", purged.getMessageCount()); - return result; - } - - private static Object listConsumers(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - String name = queueName(params); - String vhost = effectiveVhost(params, conn); - JsonElement queue = managementGet(conn, - "/api/queues/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(name)); - if (!queue.isJsonObject()) { - throw new IllegalStateException("Unexpected management API response for queue details"); - } - return Collections.singletonMap("consumers", consumersFromQueueInfo(queue.getAsJsonObject())); - } - - /** - * Map the management API queue detail's {@code consumer_details} array to the - * bridge's consumer shape. A queue without consumers may omit the array - * entirely, which maps to an empty list. - */ - static List> consumersFromQueueInfo(JsonObject info) { - List> consumers = new ArrayList<>(); - JsonElement details = info.get("consumer_details"); - if (details == null || !details.isJsonArray()) { - return consumers; - } - for (JsonElement element : details.getAsJsonArray()) { - if (!element.isJsonObject()) { - continue; - } - JsonObject consumer = element.getAsJsonObject(); - Map entry = new LinkedHashMap<>(); - String channelName = ""; - JsonElement channelDetails = consumer.get("channel_details"); - if (channelDetails != null && channelDetails.isJsonObject()) { - channelName = stringOrEmpty(channelDetails.getAsJsonObject(), "name"); - } - entry.put("name", channelName); - entry.put("tag", stringOrEmpty(consumer, "consumer_tag")); - entry.put("active", boolOrDefault(consumer, "active", false)); - entry.put("ackRequired", boolOrDefault(consumer, "ack_required", false)); - Integer prefetch = integerOrNull(consumer, "prefetch_count"); - if (prefetch != null) { - entry.put("prefetch", prefetch); - } - consumers.add(entry); - } - return consumers; - } - - // ----------------------------------------------------------------------- - // Namespaces (virtual hosts) - // ----------------------------------------------------------------------- - - private static Object listNamespaces(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - JsonElement vhosts = managementGet(conn, "/api/vhosts"); - if (!vhosts.isJsonArray()) { - throw new IllegalStateException("Unexpected management API response for vhost listing"); - } - - List> namespaces = new ArrayList<>(); - for (JsonElement element : vhosts.getAsJsonArray()) { - if (!element.isJsonObject()) { - continue; - } - namespaces.add(Collections.singletonMap("name", stringOrEmpty(element.getAsJsonObject(), "name"))); - } - return Collections.singletonMap("namespaces", namespaces); - } - - private static Object createNamespace(JsonObject params) throws Exception { - String namespace = namespaceName(params); - JsonObject conn = requireConnectionConfig(params); - managementSend(conn, "PUT", "/api/vhosts/" + urlEncodeVhost(namespace)); - return Collections.singletonMap("ok", true); - } - - private static Object deleteNamespace(JsonObject params) throws Exception { - String namespace = namespaceName(params); - // The default vhost is protected even before checking connectivity, so - // the guard error is semantic rather than a connection failure. - assertNamespaceDeletable(namespace, null); - JsonObject conn = requireConnectionConfig(params); - assertNamespaceDeletable(namespace, stringOrDefault(conn, "virtual_host", "/")); - managementSend(conn, "DELETE", "/api/vhosts/" + urlEncodeVhost(namespace)); - return Collections.singletonMap("ok", true); - } - - /** Guard rails for vhost deletion: never "/", never the vhost in use. */ - static void assertNamespaceDeletable(String namespace, String connectedVhost) { - if ("/".equals(namespace)) { - throw new IllegalArgumentException("The default virtual host '/' cannot be deleted"); - } - if (connectedVhost != null && namespace.equals(connectedVhost)) { - throw new IllegalArgumentException( - "Cannot delete the virtual host '" + namespace + "' while connected to it"); - } - } - - private static String namespaceName(JsonObject params) { - String name = stringOrEmpty(params, "namespace"); - if (name.isBlank()) { - throw new IllegalArgumentException("namespace is required"); - } - // '*' is the all-vhosts marker used by listings, never a real vhost name; - // without this guard a create/delete would address /api/vhosts/%2A. - if ("*".equals(name.trim())) { - throw new IllegalArgumentException("namespace create/delete requires a specific virtual host (all-vhosts context)"); - } - return name; - } - - // ----------------------------------------------------------------------- - // Exchanges & bindings - // ----------------------------------------------------------------------- - - private static final Set EXCHANGE_TYPES = Set.of("direct", "fanout", "topic", "headers"); - - private static Object listExchanges(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - boolean allVhosts = allVhostsRequested(params); - JsonArray exchanges = managementGetAll(conn, managementListPath(params, conn, "exchanges")); - - List> result = new ArrayList<>(); - for (JsonElement element : exchanges) { - if (!element.isJsonObject()) { - continue; - } - JsonObject exchange = element.getAsJsonObject(); - Map info = exchangeInfoFromJson(exchange); - if (allVhosts) { - attachVhost(info, exchange); - } - result.add(info); - } - result.sort(Comparator.comparing(m -> (String) m.get("name"))); - return Collections.singletonMap("exchanges", result); - } - - /** - * Map one management API exchange entry to the bridge shape. The default - * exchange ("") reports an empty type in the API; surface it as "default". - */ - static Map exchangeInfoFromJson(JsonObject exchange) { - Map info = new LinkedHashMap<>(); - info.put("name", stringOrEmpty(exchange, "name")); - String type = stringOrEmpty(exchange, "type"); - info.put("type", type.isEmpty() ? "default" : type); - info.put("durable", boolOrDefault(exchange, "durable", false)); - info.put("autoDelete", boolOrDefault(exchange, "auto_delete", false)); - info.put("internal", boolOrDefault(exchange, "internal", false)); - return info; - } - - private static Object createExchange(JsonObject params) throws Exception { - // Validate before touching connectivity so type errors are semantic. - String name = exchangeName(params); - String type = validateExchangeType(stringOrEmpty(params, "type")); - JsonObject conn = requireConnectionConfig(params); - String vhost = effectiveVhost(params, conn); - - JsonObject body = new JsonObject(); - body.addProperty("type", type); - body.addProperty("durable", boolOrDefault(params, "durable", true)); - body.addProperty("auto_delete", boolOrDefault(params, "autoDelete", false)); - managementSend(conn, "PUT", - "/api/exchanges/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(name), - body); - return Collections.singletonMap("ok", true); - } - - private static Object deleteExchange(JsonObject params) throws Exception { - // Guard before connectivity: the error is semantic, not a connection failure. - String name = stringOrEmpty(params, "name"); - assertExchangeDeletable(name); - if (name.isBlank()) { - throw new IllegalArgumentException("name is required"); - } - JsonObject conn = requireConnectionConfig(params); - String vhost = effectiveVhost(params, conn); - managementSend(conn, "DELETE", - "/api/exchanges/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(name)); - return Collections.singletonMap("ok", true); - } - - /** Exchange type whitelist; anything else is rejected before hitting the broker. */ - static String validateExchangeType(String type) { - if (!EXCHANGE_TYPES.contains(type)) { - throw new IllegalArgumentException( - "Invalid exchange type '" + type + "'. Supported types: direct, fanout, topic, headers"); - } - return type; - } - - /** Guard rails for exchange deletion: never the default exchange, never amq.* built-ins. */ - static void assertExchangeDeletable(String name) { - if (name.isEmpty()) { - throw new IllegalArgumentException("The default exchange cannot be deleted"); - } - if (name.startsWith("amq.")) { - throw new IllegalArgumentException("The built-in exchange '" + name + "' cannot be deleted"); - } - } - - private static String exchangeName(JsonObject params) { - String name = stringOrEmpty(params, "name"); - if (name.isBlank()) { - throw new IllegalArgumentException("name is required"); - } - return name; - } - - private static Object listBindings(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - boolean allVhosts = allVhostsRequested(params); - JsonArray bindings = managementGetAll(conn, managementListPath(params, conn, "bindings")); - String exchange = stringOrEmpty(params, "exchange"); - String queue = stringOrEmpty(params, "queue"); - - List> result = new ArrayList<>(); - for (JsonElement element : bindings) { - if (!element.isJsonObject()) { - continue; - } - Map binding = bindingInfoFromJson(element.getAsJsonObject()); - if (!exchange.isEmpty() && !exchange.equals(binding.get("source"))) { - continue; - } - // A queue filter means "bindings feeding this queue": the - // destination must be the queue itself, not an exchange. - if (!queue.isEmpty() && !(queue.equals(binding.get("destination")) - && "queue".equals(binding.get("destinationType")))) { - continue; - } - if (allVhosts) { - attachVhost(binding, element.getAsJsonObject()); - } - result.add(binding); - } - return Collections.singletonMap("bindings", result); - } - - /** Map one management API binding entry (snake_case) to the bridge shape (camelCase). */ - static Map bindingInfoFromJson(JsonObject binding) { - Map info = new LinkedHashMap<>(); - info.put("source", stringOrEmpty(binding, "source")); - info.put("destination", stringOrEmpty(binding, "destination")); - info.put("destinationType", stringOrEmpty(binding, "destination_type")); - info.put("routingKey", stringOrEmpty(binding, "routing_key")); - JsonElement arguments = binding.get("arguments"); - if (arguments != null && arguments.isJsonObject() && !arguments.getAsJsonObject().isEmpty()) { - Map args = new LinkedHashMap<>(); - for (Map.Entry entry : arguments.getAsJsonObject().entrySet()) { - if (entry.getValue().isJsonNull()) { - continue; - } - Object value = argumentValue(entry.getValue()); - args.put(entry.getKey(), value != null ? value : entry.getValue().toString()); - } - info.put("arguments", args); - } - return info; - } - - private static Object bind(JsonObject params) throws Exception { - applyBinding(params, true); - return Collections.singletonMap("ok", true); - } - - private static Object unbind(JsonObject params) throws Exception { - applyBinding(params, false); - return Collections.singletonMap("ok", true); - } - - /** - * Bind or unbind via AMQP. Queue destinations use queueBind/queueUnbind; - * exchange destinations (exchange-to-exchange) use exchangeBind/exchangeUnbind. - */ - private static void applyBinding(JsonObject params, boolean bind) throws Exception { - String source = requireBindingName(params, "source"); - String destination = requireBindingName(params, "destination"); - String destinationType = stringOrDefault(params, "destinationType", - stringOrEmpty(params, "destination_type")); - // Validate the destination type before touching connectivity so a bad - // value fails fast instead of surfacing as a connection error. - if (!"queue".equals(destinationType) && !"exchange".equals(destinationType)) { - throw new IllegalArgumentException( - "destinationType must be 'queue' or 'exchange', got '" + destinationType + "'"); - } - String routingKey = stringOrDefault(params, "routingKey", stringOrEmpty(params, "routing_key")); - Map arguments = bindingArguments(params); - Channel ch = channelFor(params); - - switch (destinationType) { - case "queue" -> { - if (bind) { - ch.queueBind(destination, source, routingKey, arguments); - } else { - ch.queueUnbind(destination, source, routingKey, arguments); - } - } - case "exchange" -> { - if (bind) { - ch.exchangeBind(destination, source, routingKey, arguments); - } else { - ch.exchangeUnbind(destination, source, routingKey, arguments); - } - } - default -> throw new IllegalStateException("unreachable"); - } - } - - private static String requireBindingName(JsonObject params, String key) { - String name = stringOrEmpty(params, key); - if (name.isBlank()) { - throw new IllegalArgumentException(key + " is required"); - } - return name; - } - - private static Map bindingArguments(JsonObject params) { - Map arguments = new HashMap<>(); - JsonObject args = params.has("arguments") && params.get("arguments").isJsonObject() - ? params.getAsJsonObject("arguments") : null; - if (args != null) { - for (Map.Entry entry : args.entrySet()) { - Object value = argumentValue(entry.getValue()); - if (value != null) { - arguments.put(entry.getKey(), value); - } - } - } - return arguments; - } - - // ----------------------------------------------------------------------- - // Client connections & channels - // ----------------------------------------------------------------------- - - private static Object listClientConnections(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - JsonArray connections = managementGetAll(conn, "/api/connections"); - boolean allVhosts = allVhostsRequested(params); - String vhostFilter = vhostFilter(params, conn); - - List> result = new ArrayList<>(); - for (JsonElement element : connections) { - if (!element.isJsonObject()) { - continue; - } - JsonObject connection = element.getAsJsonObject(); - if (!vhostFilter.isEmpty() && !vhostFilter.equals(stringOrEmpty(connection, "vhost"))) { - continue; - } - Map info = clientConnectionInfoFromJson(connection); - if (allVhosts) { - attachVhost(info, connection); - } - result.add(info); - } - result.sort(Comparator.comparing(m -> (String) m.get("name"))); - return Collections.singletonMap("connections", result); - } - - /** - * Map one management API connection entry (snake_case) to the bridge shape - * (camelCase). Rates come from the *_oct_details blocks; connected_at is a - * millisecond timestamp. Both are omitted when the broker does not report them. - */ - static Map clientConnectionInfoFromJson(JsonObject connection) { - Map info = new LinkedHashMap<>(); - info.put("name", stringOrEmpty(connection, "name")); - info.put("user", stringOrEmpty(connection, "user")); - info.put("peerHost", stringOrEmpty(connection, "peer_host")); - info.put("peerPort", longOrDefault(connection, "peer_port", 0)); - info.put("state", stringOrEmpty(connection, "state")); - info.put("channels", longOrDefault(connection, "channels", 0)); - Double recvRate = rateFromDetails(connection, "recv_oct_details"); - if (recvRate != null) { - info.put("recvRate", recvRate); - } - Double sendRate = rateFromDetails(connection, "send_oct_details"); - if (sendRate != null) { - info.put("sendRate", sendRate); - } - Long connectedAt = longOrNull(connection, "connected_at"); - if (connectedAt != null) { - info.put("connectedAt", connectedAt); - } - return info; - } - - /** Per-second byte rate from a {@code recv_oct_details}/{@code send_oct_details} block. */ - static Double rateFromDetails(JsonObject object, String key) { - JsonElement details = object.get(key); - if (details == null || !details.isJsonObject()) { - return null; - } - JsonElement rate = details.getAsJsonObject().get("rate"); - return rate == null || rate.isJsonNull() ? null : rate.getAsDouble(); - } - - private static Object listClientChannels(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - JsonArray channels = managementGetAll(conn, "/api/channels"); - String connectionFilter = stringOrEmpty(params, "connection"); - boolean allVhosts = allVhostsRequested(params); - String vhostFilter = vhostFilter(params, conn); - - List> result = new ArrayList<>(); - for (JsonElement element : channels) { - if (!element.isJsonObject()) { - continue; - } - JsonObject channel = element.getAsJsonObject(); - if (!vhostFilter.isEmpty() && !vhostFilter.equals(stringOrEmpty(channel, "vhost"))) { - continue; - } - Map info = channelInfoFromJson(channel); - if (allVhosts) { - attachVhost(info, channel); - } - if (!connectionFilter.isEmpty() && !channelMatchesConnection(info, connectionFilter)) { - continue; - } - result.add(info); - } - result.sort(Comparator.comparing(m -> (String) m.get("name"))); - return Collections.singletonMap("channels", result); - } - - /** Map one management API channel entry (snake_case) to the bridge shape (camelCase). */ - static Map channelInfoFromJson(JsonObject channel) { - Map info = new LinkedHashMap<>(); - info.put("name", stringOrEmpty(channel, "name")); - JsonElement connectionDetails = channel.get("connection_details"); - if (connectionDetails != null && connectionDetails.isJsonObject()) { - String connectionName = stringOrEmpty(connectionDetails.getAsJsonObject(), "name"); - if (!connectionName.isEmpty()) { - info.put("connectionName", connectionName); - } - } - info.put("state", stringOrEmpty(channel, "state")); - Integer prefetch = integerOrNull(channel, "prefetch_count"); - if (prefetch != null) { - info.put("prefetch", prefetch); - } - Long unacked = longOrNull(channel, "messages_unacknowledged"); - if (unacked != null) { - info.put("messagesUnacked", unacked); - } - Long consumers = longOrNull(channel, "consumer_count"); - if (consumers != null) { - info.put("consumerCount", consumers); - } - return info; - } - - /** - * A channel belongs to a connection when its {@code connection_details.name} - * matches, or when its own name starts with the connection name (channel - * names are "{connectionName} ({channelNumber})"). - */ - static boolean channelMatchesConnection(Map channelInfo, String connectionName) { - if (connectionName.equals(channelInfo.get("connectionName"))) { - return true; - } - Object name = channelInfo.get("name"); - return name instanceof String && ((String) name).startsWith(connectionName); - } - - private static Object closeClientConnection(JsonObject params) throws Exception { - // Validate before touching connectivity so the error is semantic. - String name = stringOrEmpty(params, "name"); - if (name.isBlank()) { - throw new IllegalArgumentException("name is required"); - } - JsonObject conn = requireConnectionConfig(params); - managementSend(conn, "DELETE", "/api/connections/" + urlEncodeName(name)); - return Collections.singletonMap("ok", true); - } - - /** - * URL-encode a connection name for the management API path. Connection names - * contain " -> " and spaces, so URLEncoder's '+' for spaces must become %20. - */ - static String urlEncodeName(String name) { - return urlEncodePathSegment(name); - } - - // ----------------------------------------------------------------------- - // Users & permissions - // ----------------------------------------------------------------------- - - private static Object listUsers(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - JsonArray users = managementGetAll(conn, "/api/users"); - - List> result = new ArrayList<>(); - for (JsonElement element : users) { - if (!element.isJsonObject()) { - continue; - } - result.add(userInfoFromJson(element.getAsJsonObject())); - } - result.sort(Comparator.comparing(m -> (String) m.get("name"))); - return Collections.singletonMap("users", result); - } - - /** - * Map one management API user entry to the bridge shape. The API reports tags - * as a single comma-separated string; the bridge shape carries them as an array. - */ - static Map userInfoFromJson(JsonObject user) { - Map info = new LinkedHashMap<>(); - info.put("name", stringOrEmpty(user, "name")); - info.put("tags", parseUserTags(stringOrEmpty(user, "tags"))); - return info; - } - - /** Split the management API's comma-separated tag string; blank entries are dropped. */ - static List parseUserTags(String tags) { - List result = new ArrayList<>(); - for (String tag : tags.split(",")) { - String trimmed = tag.trim(); - if (!trimmed.isEmpty()) { - result.add(trimmed); - } - } - return result; - } - - private static Object createUser(JsonObject params) throws Exception { - // Validate before touching connectivity so the errors are semantic. - String name = userName(params); - String password = stringOrEmpty(params, "password"); - if (password.isEmpty()) { - throw new IllegalArgumentException("password is required"); - } - JsonObject conn = requireConnectionConfig(params); - // PUT /api/users upserts, so "creating" the connected user would actually - // change its credentials; reject it just like deletion. - assertNotConnectedUser("create or modify", name, stringOrDefault(conn, "username", "guest")); - - JsonObject body = new JsonObject(); - body.addProperty("password", password); - body.addProperty("tags", userTagsParam(params)); - managementSend(conn, "PUT", "/api/users/" + urlEncodePathSegment(name), body); - return Collections.singletonMap("ok", true); - } - - /** Tags for user creation: accepts a JSON array or a comma-separated string. */ - static String userTagsParam(JsonObject params) { - JsonElement tags = params.get("tags"); - if (tags == null || tags.isJsonNull()) { - return ""; - } - if (tags.isJsonArray()) { - List parts = new ArrayList<>(); - for (JsonElement element : tags.getAsJsonArray()) { - String tag = element.getAsString().trim(); - if (!tag.isEmpty()) { - parts.add(tag); - } - } - return String.join(",", parts); - } - return tags.getAsString(); - } - - private static Object deleteUser(JsonObject params) throws Exception { - String name = userName(params); - JsonObject conn = requireConnectionConfig(params); - assertNotConnectedUser("delete", name, stringOrDefault(conn, "username", "guest")); - managementSend(conn, "DELETE", "/api/users/" + urlEncodePathSegment(name)); - return Collections.singletonMap("ok", true); - } - - /** Guard rail for user changes: never touch the user the agent itself connects as. */ - static void assertNotConnectedUser(String action, String name, String connectedUser) { - if (name.equals(connectedUser)) { - throw new IllegalArgumentException( - "Cannot " + action + " user '" + name + "' while connected as that user"); - } - } - - private static Object listPermissions(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - JsonElement permissions = managementGet(conn, "/api/permissions"); - if (!permissions.isJsonArray()) { - throw new IllegalStateException("Unexpected management API response for permission listing"); - } - // The management API only lists permissions cluster-wide; virtual_host and - // user are client-side filters. all_vhosts simply disables the vhost filter. - String vhostFilter = allVhostsRequested(params) ? "" : stringOrEmpty(params, "virtual_host"); - String userFilter = stringOrEmpty(params, "user"); - - List> result = new ArrayList<>(); - for (JsonElement element : permissions.getAsJsonArray()) { - if (!element.isJsonObject()) { - continue; - } - Map permission = permissionInfoFromJson(element.getAsJsonObject()); - if (!vhostFilter.isEmpty() && !vhostFilter.equals(permission.get("vhost"))) { - continue; - } - if (!userFilter.isEmpty() && !userFilter.equals(permission.get("user"))) { - continue; - } - result.add(permission); - } - result.sort(Comparator.comparing((Map m) -> (String) m.get("user")) - .thenComparing(m -> (String) m.get("vhost"))); - return Collections.singletonMap("permissions", result); - } - - /** Map one management API permission entry (user x vhost regex triple) to the bridge shape. */ - static Map permissionInfoFromJson(JsonObject permission) { - Map info = new LinkedHashMap<>(); - info.put("user", stringOrEmpty(permission, "user")); - info.put("vhost", stringOrEmpty(permission, "vhost")); - info.put("configure", stringOrEmpty(permission, "configure")); - info.put("write", stringOrEmpty(permission, "write")); - info.put("read", stringOrEmpty(permission, "read")); - return info; - } - - private static final String DEFAULT_PERMISSION_PATTERN = ".*"; - - private static Object grantPermission(JsonObject params) throws Exception { - // Validate before touching connectivity so the errors are semantic. - String user = userName(params); - String vhost = permissionVhost(params); - JsonObject conn = requireConnectionConfig(params); - - JsonObject body = new JsonObject(); - body.addProperty("configure", permissionPattern(params, "configure")); - body.addProperty("write", permissionPattern(params, "write")); - body.addProperty("read", permissionPattern(params, "read")); - managementSend(conn, "PUT", - "/api/permissions/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(user), body); - return Collections.singletonMap("ok", true); - } - - private static Object revokePermission(JsonObject params) throws Exception { - String user = userName(params); - String vhost = permissionVhost(params); - JsonObject conn = requireConnectionConfig(params); - managementSend(conn, "DELETE", - "/api/permissions/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(user)); - return Collections.singletonMap("ok", true); - } - - /** Permission pattern defaulting to ".*" (full access) when omitted or blank. */ - static String permissionPattern(JsonObject params, String key) { - String pattern = stringOrEmpty(params, key); - return pattern.isBlank() ? DEFAULT_PERMISSION_PATTERN : pattern; - } - - /** - * Vhost for permission and policy writes: the {@code *} all-vhosts sentinel - * only makes sense for listings; a write always targets one concrete vhost. - */ - static String permissionVhost(JsonObject params) { - String vhost = stringOrEmpty(params, "virtual_host"); - if (vhost.isBlank()) { - throw new IllegalArgumentException("virtual_host is required"); - } - if ("*".equals(vhost)) { - throw new IllegalArgumentException("all_vhosts is only supported for list operations"); - } - return vhost; - } - - /** User name: create/delete send {@code name}, grant/revoke send {@code user}. */ - private static String userName(JsonObject params) { - String name = stringOrEmpty(params, "name"); - if (name.isBlank()) { - name = stringOrEmpty(params, "user"); - } - if (name.isBlank()) { - throw new IllegalArgumentException("user name is required"); - } - return name; - } - - /** - * URL-encode one path segment (queue/exchange name) for the management API. - * URLEncoder is form-oriented and encodes spaces as '+', which the management - * API does not decode back in path segments (causing 404s), so '+' becomes %20. - */ - static String urlEncodePathSegment(String name) { - return URLEncoder.encode(name, StandardCharsets.UTF_8).replace("+", "%20"); - } - - // ----------------------------------------------------------------------- - // Policies - // ----------------------------------------------------------------------- - - private static Object listPolicies(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - // Accept the '*' all-vhosts sentinel as a synonym for all_vhosts=true; - // both select the vhost-less management API variant. - boolean allVhosts = allVhostsRequested(params) || "*".equals(stringOrEmpty(params, "virtual_host")); - JsonArray policies = managementGetAll(conn, allVhosts ? "/api/policies" - : "/api/policies/" + urlEncodeVhost(effectiveVhost(params, conn))); - - List> result = new ArrayList<>(); - for (JsonElement element : policies) { - if (!element.isJsonObject()) { - continue; - } - result.add(policyInfoFromJson(element.getAsJsonObject())); - } - result.sort(Comparator.comparing((Map m) -> (String) m.get("vhost")) - .thenComparing(m -> (String) m.get("name"))); - return Collections.singletonMap("policies", result); - } - - /** - * Map one management API policy entry to the bridge shape: kebab-case - * {@code apply-to} becomes camelCase {@code applyTo}, and the definition map - * is passed through with plain values. Each policy always carries its own - * {@code vhost}, so flat and cross-vhost listings share one shape. - */ - static Map policyInfoFromJson(JsonObject policy) { - Map info = new LinkedHashMap<>(); - info.put("name", stringOrEmpty(policy, "name")); - info.put("vhost", stringOrEmpty(policy, "vhost")); - info.put("pattern", stringOrEmpty(policy, "pattern")); - info.put("applyTo", stringOrEmpty(policy, "apply-to")); - info.put("priority", longOrDefault(policy, "priority", 0)); - Map definition = new LinkedHashMap<>(); - JsonElement rawDefinition = policy.get("definition"); - if (rawDefinition != null && rawDefinition.isJsonObject()) { - for (Map.Entry entry : rawDefinition.getAsJsonObject().entrySet()) { - if (entry.getValue().isJsonNull()) { - continue; - } - Object value = argumentValue(entry.getValue()); - definition.put(entry.getKey(), value != null ? value : entry.getValue().toString()); - } - } - info.put("definition", definition); - return info; - } - - private static Object setPolicy(JsonObject params) throws Exception { - // Validate before touching connectivity so the errors are semantic. - String vhost = permissionVhost(params); - String name = policyName(params); - String pattern = stringOrEmpty(params, "pattern"); - if (pattern.isBlank()) { - throw new IllegalArgumentException("pattern is required"); - } - JsonElement definition = params.get("definition"); - if (definition == null || !definition.isJsonObject()) { - throw new IllegalArgumentException("definition is required"); - } - JsonObject conn = requireConnectionConfig(params); - - JsonObject body = new JsonObject(); - body.addProperty("pattern", pattern); - // The bridge sends camelCase applyTo; the management API wants apply-to. - // applyTo defaults to queues and priority to 0, matching broker defaults. - body.addProperty("apply-to", stringOrDefault(params, "applyTo", "queues")); - body.addProperty("priority", intOrDefault(params, "priority", 0)); - body.add("definition", definition); - managementSend(conn, "PUT", - "/api/policies/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(name), body); - return Collections.singletonMap("ok", true); - } - - private static Object deletePolicy(JsonObject params) throws Exception { - // Validate before touching connectivity so the errors are semantic. - String vhost = permissionVhost(params); - String name = policyName(params); - JsonObject conn = requireConnectionConfig(params); - managementSend(conn, "DELETE", - "/api/policies/" + urlEncodeVhost(vhost) + "/" + urlEncodePathSegment(name)); - return Collections.singletonMap("ok", true); - } - - /** Policy name for set/delete. */ - private static String policyName(JsonObject params) { - String name = stringOrEmpty(params, "name"); - if (name.isBlank()) { - throw new IllegalArgumentException("name is required"); - } - return name; - } - - // ----------------------------------------------------------------------- - // Messages - // ----------------------------------------------------------------------- - - private static Object peekMessages(JsonObject params) throws Exception { - String queue = queueName(params); - long offset = normalizePeekOffset(longOrDefault(params, "offset", 0)); - int count = normalizePeekCount(intOrDefault(params, "count", 10)); - long totalToFetch = Math.min(offset + count, MAX_PEEK_MESSAGES); - if (offset >= MAX_PEEK_MESSAGES) { - return Collections.singletonMap("messages", Collections.emptyList()); - } - - Channel ch; - Connection ownedConnection = null; - if (cachedConnection != null) { - ch = channelFor(params); - } else { - // Not connected: open a short-lived connection from inline params, - // mirroring how the Kafka agent accepts a `connection` object for peek. - JsonObject conn = connectionObject(params); - String vhost = stringOrNull(params, "virtual_host"); - if (vhost != null && !vhost.isBlank()) { - conn = conn.deepCopy(); - conn.addProperty("virtual_host", vhost); - } - ownedConnection = openConnection(conn); - ch = ownedConnection.createChannel(); - } - try { - List fetched = new ArrayList<>(); - long lastDeliveryTag = -1; - for (long i = 0; i < totalToFetch; i++) { - GetResponse response = ch.basicGet(queue, false); - if (response == null) { - break; - } - fetched.add(response); - lastDeliveryTag = response.getEnvelope().getDeliveryTag(); - } - // Requeue everything so peeking never consumes messages. - if (lastDeliveryTag >= 0) { - ch.basicNack(lastDeliveryTag, true, true); - } - - List> messages = new ArrayList<>(); - for (long i = offset; i < fetched.size() && messages.size() < count; i++) { - messages.add(peekedMessageFromGetResponse(queue, i, fetched.get((int) i))); - } - return Collections.singletonMap("messages", messages); - } finally { - if (ownedConnection != null) { - closeQuietly(ch); - closeQuietly(ownedConnection); - } - } - } - - static long normalizePeekOffset(long requestedOffset) { - return Math.max(0, requestedOffset); - } - - /** - * Routing key for publishes: {@code routing_key}/{@code routingKey} win; - * the Rust bridge sends the message key as {@code key}; the default is the - * queue name so publishes through the default exchange reach the queue. - */ - static String resolveRoutingKey(JsonObject params, String queue) { - String routingKey = stringOrDefault(params, "routing_key", ""); - if (routingKey.isEmpty()) { - routingKey = stringOrDefault(params, "routingKey", ""); - } - if (routingKey.isEmpty()) { - routingKey = stringOrDefault(params, "key", ""); - } - // A blank key must not win over the queue fallback: publishing through - // the default exchange with an empty routing key silently drops the message. - if (routingKey.isBlank()) { - routingKey = queue; - } - return routingKey; - } - - static int normalizePeekCount(int requestedCount) { - return Math.max(1, requestedCount); - } - - private static Map peekedMessageFromGetResponse(String queue, long index, GetResponse response) { - Map msg = new LinkedHashMap<>(); - msg.put("topic", queue); - msg.put("offset", index); - msg.put("exchange", response.getEnvelope().getExchange()); - msg.put("routingKey", response.getEnvelope().getRoutingKey()); - msg.put("redelivered", response.getEnvelope().isRedeliver()); - msg.put("deliveryTag", response.getEnvelope().getDeliveryTag()); - - AMQP.BasicProperties props = response.getProps(); - if (props != null && props.getMessageId() != null) { - msg.put("messageId", props.getMessageId()); - } - Date timestamp = props != null ? props.getTimestamp() : null; - msg.put("timestamp", timestamp != null ? timestamp.getTime() : 0L); - - Map headers = new LinkedHashMap<>(); - if (props != null && props.getHeaders() != null) { - for (Map.Entry entry : props.getHeaders().entrySet()) { - headers.put(entry.getKey(), String.valueOf(entry.getValue())); - } - } - msg.put("headers", headers); - - byte[] body = response.getBody(); - if (body != null) { - msg.put("payloadBase64", Base64.getEncoder().encodeToString(body)); - String text = tryDecodeUtf8(body); - if (text != null) { - msg.put("payloadText", text); - } - } else { - msg.put("payloadBase64", ""); - } - return msg; - } - - private static Object sendMessage(JsonObject params) throws Exception { - Channel ch = channelFor(params); - String queue = queueName(params); - String exchange = stringOrDefault(params, "exchange", ""); - String routingKey = resolveRoutingKey(params, queue); - - String payloadBase64 = stringOrEmpty(params, "payloadBase64"); - byte[] body = payloadBase64.isEmpty() ? new byte[0] : Base64.getDecoder().decode(payloadBase64); - - AMQP.BasicProperties properties = null; - JsonObject headers = params.has("headers") && params.get("headers").isJsonObject() - ? params.getAsJsonObject("headers") : null; - if (headers != null) { - Map headerMap = new HashMap<>(); - for (Map.Entry entry : headers.entrySet()) { - Object value = argumentValue(entry.getValue()); - if (value != null) { - headerMap.put(entry.getKey(), value); - } - } - properties = new AMQP.BasicProperties.Builder().headers(headerMap).build(); - } - - ch.basicPublish(exchange, routingKey, properties, body); - - Map result = new LinkedHashMap<>(); - result.put("ok", true); - result.put("exchange", exchange); - result.put("routingKey", routingKey); - return result; - } - - // ----------------------------------------------------------------------- - // Cluster / monitoring - // ----------------------------------------------------------------------- - - private static Object describeCluster(JsonObject params) throws Exception { - Connection conn = requireConnection(); - JsonObject connConfig = currentConnectionConfig(params); - Map serverProps = conn.getServerProperties(); - - List> nodes = new ArrayList<>(); - if (connConfig != null) { - for (Address address : resolveAddresses(connConfig)) { - Map node = new LinkedHashMap<>(); - node.put("name", address.getHost()); - node.put("port", address.getPort()); - nodes.add(node); - } - } - - Map result = new LinkedHashMap<>(); - result.put("clusterName", serverString(serverProps, "cluster_name")); - result.put("product", serverString(serverProps, "product")); - result.put("version", serverString(serverProps, "version")); - result.put("platform", serverString(serverProps, "platform")); - result.put("nodes", nodes); - result.put("nodeCount", nodes.size()); - return result; - } - - private static Object getOverview(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - JsonElement overview = managementGet(conn, "/api/overview"); - if (!overview.isJsonObject()) { - throw new IllegalStateException("Unexpected management API response for cluster overview"); - } - return overviewInfoFromJson(overview.getAsJsonObject()); - } - - /** - * Map the management API overview (snake_case totals and message stats) to - * the bridge shape (camelCase). Rates come from each stat's - * {@code *_details.rate} block; anything the broker does not report is - * omitted rather than zeroed. - */ - static Map overviewInfoFromJson(JsonObject overview) { - Map info = new LinkedHashMap<>(); - putIfPresent(info, "messagesReady", nestedLongOrNull(overview, "queue_totals", "messages_ready")); - putIfPresent(info, "messagesUnacked", nestedLongOrNull(overview, "queue_totals", "messages_unacknowledged")); - - JsonElement stats = overview.get("message_stats"); - if (stats != null && stats.isJsonObject()) { - JsonObject messageStats = stats.getAsJsonObject(); - putIfPresent(info, "publishRate", rateFromDetails(messageStats, "publish_details")); - putIfPresent(info, "deliverRate", rateFromDetails(messageStats, "deliver_get_details")); - putIfPresent(info, "ackRate", rateFromDetails(messageStats, "ack_details")); - } - - putIfPresent(info, "totalQueues", nestedLongOrNull(overview, "object_totals", "queues")); - putIfPresent(info, "totalExchanges", nestedLongOrNull(overview, "object_totals", "exchanges")); - putIfPresent(info, "totalConnections", nestedLongOrNull(overview, "object_totals", "connections")); - putIfPresent(info, "totalChannels", nestedLongOrNull(overview, "object_totals", "channels")); - putIfPresent(info, "totalConsumers", nestedLongOrNull(overview, "object_totals", "consumers")); - return info; - } - - private static Object listNodes(JsonObject params) throws Exception { - JsonObject conn = requireConnectionConfig(params); - JsonElement nodes = managementGet(conn, "/api/nodes"); - if (!nodes.isJsonArray()) { - throw new IllegalStateException("Unexpected management API response for node listing"); - } - - List> result = new ArrayList<>(); - for (JsonElement element : nodes.getAsJsonArray()) { - if (!element.isJsonObject()) { - continue; - } - result.add(nodeInfoFromJson(element.getAsJsonObject())); - } - result.sort(Comparator.comparing(m -> (String) m.get("name"))); - return Collections.singletonMap("nodes", result); - } - - /** - * Map one management API node entry (snake_case) to the bridge shape - * (camelCase). The API reports uptime in milliseconds; resource counters - * the broker does not report are omitted rather than zeroed. - */ - static Map nodeInfoFromJson(JsonObject node) { - Map info = new LinkedHashMap<>(); - info.put("name", stringOrEmpty(node, "name")); - info.put("running", boolOrDefault(node, "running", false)); - putIfPresent(info, "memUsed", longOrNull(node, "mem_used")); - putIfPresent(info, "memLimit", longOrNull(node, "mem_limit")); - putIfPresent(info, "diskFree", longOrNull(node, "disk_free")); - putIfPresent(info, "fdUsed", longOrNull(node, "fd_used")); - putIfPresent(info, "fdTotal", longOrNull(node, "fd_total")); - putIfPresent(info, "socketsUsed", longOrNull(node, "sockets_used")); - putIfPresent(info, "socketsTotal", longOrNull(node, "sockets_total")); - putIfPresent(info, "uptimeMs", longOrNull(node, "uptime")); - return info; - } - - /** Long value one level down (e.g. {@code object_totals.queues}); null when absent. */ - static Long nestedLongOrNull(JsonObject object, String block, String key) { - JsonElement element = object.get(block); - if (element == null || !element.isJsonObject()) { - return null; - } - return longOrNull(element.getAsJsonObject(), key); - } - - /** Adds the value only when the broker reported it (missing stats stay absent). */ - private static void putIfPresent(Map info, String key, Object value) { - if (value != null) { - info.put(key, value); - } - } - - // ----------------------------------------------------------------------- - // HTTP management API helpers - // ----------------------------------------------------------------------- - - static JsonElement managementGet(JsonObject conn, String path) throws Exception { - return managementRequest(conn, "GET", path); - } - - /** Management API call without a JSON body (PUT/DELETE); accepts 2xx. */ - static JsonElement managementSend(JsonObject conn, String method, String path) throws Exception { - return managementRequest(conn, method, path); - } - - /** Management API call with a JSON body (PUT/POST); accepts 2xx. */ - static JsonElement managementSend(JsonObject conn, String method, String path, JsonObject body) throws Exception { - return managementRequest(conn, method, path, body); - } - - static JsonElement managementRequest(JsonObject conn, String method, String path) throws Exception { - return managementRequest(conn, method, path, null); - } - - static JsonElement managementRequest(JsonObject conn, String method, String path, JsonObject body) throws Exception { - // Candidates are tried in order; only connection-level failures - // (refused/timeout/DNS) move to the next candidate. A non-2xx HTTP - // status means the endpoint answered, so the answer is final. - IOException lastConnectionError = null; - for (String baseUrl : managementBaseUrls(conn)) { - try { - return managementRequestOnce(baseUrl, conn, method, path, body); - } catch (IOException e) { - lastConnectionError = e; - } - } - throw lastConnectionError != null ? lastConnectionError - : new IllegalStateException("No management API endpoint candidates"); - } - - private static JsonElement managementRequestOnce(String baseUrl, JsonObject conn, - String method, String path, JsonObject body) throws Exception { - URL url = URI.create(baseUrl + path).toURL(); - HttpURLConnection http = (HttpURLConnection) url.openConnection(); - try { - // tls_skip_verify previously only applied to AMQP; honor it for the - // management API too, or self-signed brokers fail every HTTP call. - if (tlsSkipVerify(conn) && http instanceof HttpsURLConnection https) { - https.setSSLSocketFactory(trustAllSslContext().getSocketFactory()); - https.setHostnameVerifier((hostname, session) -> true); - } - http.setRequestMethod(method); - http.setConnectTimeout(10_000); - http.setReadTimeout(10_000); - http.setRequestProperty("Authorization", - basicAuthHeader(credentialOrGuest(conn, "username"), - credentialOrGuest(conn, "password"))); - if (body != null) { - http.setDoOutput(true); - http.setRequestProperty("Content-Type", "application/json"); - try (OutputStream out = http.getOutputStream()) { - out.write(GSON.toJson(body).getBytes(StandardCharsets.UTF_8)); - } - } - int status = http.getResponseCode(); - if (status < 200 || status >= 300) { - throw new IllegalStateException(managementErrorMessage(status, method, path)); - } - if (status == 204) { - return JsonNull.INSTANCE; - } - try (InputStream in = http.getInputStream()) { - String responseBody = new String(in.readAllBytes(), StandardCharsets.UTF_8); - if (responseBody.isBlank()) { - return JsonNull.INSTANCE; - } - return JsonParser.parseString(responseBody); - } - } finally { - http.disconnect(); - } - } - - static String managementBaseUrl(String host, int port, boolean tls) { - return (tls ? "https" : "http") + "://" + host + ":" + port; - } - - /** - * Candidate management API base URLs. An explicit {@code management_url} - * wins and is used verbatim (scheme/host/port/path prefix, e.g. a reverse - * proxy mount like {@code https://proxy:8443/rmq}); otherwise one candidate - * per AMQP address is derived with the management port, and - * {@link #managementRequest} fails over across them. - */ - static List managementBaseUrls(JsonObject conn) { - String explicit = stringOrNull(conn, "management_url"); - if (explicit != null && !explicit.isBlank()) { - return List.of(normalizeManagementUrl(explicit)); - } - boolean tls = managementTls(conn); - int port = managementPort(conn, tls); - List baseUrls = new ArrayList<>(); - for (Address address : resolveAddresses(conn)) { - baseUrls.add(managementBaseUrl(address.getHost(), port, tls)); - } - return baseUrls; - } - - /** - * Trailing slashes are trimmed so base + "/api/..." joins cleanly; the path - * prefix itself is kept verbatim (no re-encoding). - */ - static String normalizeManagementUrl(String url) { - String trimmed = url.trim(); - while (trimmed.endsWith("/")) { - trimmed = trimmed.substring(0, trimmed.length() - 1); - } - return trimmed; - } - - /** - * Whether the derived management endpoint uses TLS. Only explicit tls/ssl - * parameters count: tls_skip_verify is a verification flag, not a scheme - * indicator, and must not flip the management API to https. - */ - static boolean managementTls(JsonObject conn) { - return (conn.has("tls") && conn.get("tls").isJsonObject()) - || boolProperty(conn, "ssl") - || boolProperty(conn, "tls"); - } - - /** - * Username/password with blank normalization: a missing, null, or - * whitespace-only credential falls back to "guest". Without this an empty - * string from the bridge authenticates as ":" and fails confusingly. - */ - static String credentialOrGuest(JsonObject conn, String key) { - String value = stringOrNull(conn, key); - return value == null || value.isBlank() ? "guest" : value; - } - - private static final int MANAGEMENT_PAGE_SIZE = 100; - - /** - * Fetch every item of a management API list endpoint. RabbitMQ answers a - * paginated request ({@code page}/{@code page_size}) with - * {@code {items, page, page_count, total_count}}, so the loop walks to the - * last page; brokers that ignore the parameters answer with a plain array, - * which is returned as-is. - */ - static JsonArray managementGetAll(JsonObject conn, String path) throws Exception { - JsonArray all = new JsonArray(); - for (int page = 1;; page++) { - String separator = path.contains("?") ? "&" : "?"; - JsonElement response = managementGet(conn, - path + separator + "page=" + page + "&page_size=" + MANAGEMENT_PAGE_SIZE); - if (response.isJsonArray()) { - response.getAsJsonArray().forEach(all::add); - return all; - } - if (!response.isJsonObject() || !response.getAsJsonObject().has("items")) { - throw new IllegalStateException( - "Unexpected management API response for list endpoint " + path); - } - JsonObject paged = response.getAsJsonObject(); - JsonElement items = paged.get("items"); - if (items.isJsonArray()) { - items.getAsJsonArray().forEach(all::add); - } - Integer pageCount = integerOrNull(paged, "page_count"); - if (pageCount == null || page >= pageCount) { - return all; - } - } - } - - /** - * Error text for a non-2xx management API response. 401/403 mean the plugin - * answered but rejected the credentials or the user's management tag, so - * blaming the plugin would mislead debugging; other statuses keep the - * plugin hint (connection refused/timeouts never reach this method). - */ - static String managementErrorMessage(int status, String method, String path) { - String base = "RabbitMQ management API returned HTTP " + status + " for " + method + " " + path + "."; - if (status == 401 || status == 403) { - return base + " Hint: check the username/password and that the user has a management" - + " permission tag (management, policymaker, monitoring, or administrator)."; - } - return base + " The rabbitmq_management plugin must be enabled for this operation."; - } - - static int managementPort(JsonObject conn, boolean tls) { - Integer configured = null; - JsonObject properties = conn.has("properties") && conn.get("properties").isJsonObject() - ? conn.getAsJsonObject("properties") : null; - if (properties != null) { - configured = integerProperty(properties, "management_port"); - } - if (configured != null) { - return configured; - } - return tls ? DEFAULT_MANAGEMENT_TLS_PORT : DEFAULT_MANAGEMENT_PORT; - } - - static String basicAuthHeader(String username, String password) { - String credentials = username + ":" + password; - return "Basic " + Base64.getEncoder().encodeToString(credentials.getBytes(StandardCharsets.UTF_8)); - } - - static String urlEncodeVhost(String vhost) { - return URLEncoder.encode(vhost, StandardCharsets.UTF_8).replace("+", "%20"); - } - - // ----------------------------------------------------------------------- - // Helpers - // ----------------------------------------------------------------------- - - private static final Pattern QUOTED_NAME = Pattern.compile("'([^']+)'"); - private static final Pattern DECLARED_RESOURCE_NAME = - Pattern.compile("for (queue|exchange) '([^']+)'"); - - static String normalizeErrorMessage(Exception e) { - // AMQP channel/connection shutdowns carry a broker reply code; map the - // known ones to a readable message instead of leaking raw AMQP text. - String friendly = amqpFriendlyMessage(e); - if (friendly != null) { - return friendly; - } - String message = e.getMessage() == null || e.getMessage().isBlank() - ? e.getClass().getName() - : e.getMessage(); - Throwable root = rootCause(e); - if (root != e && root.getMessage() != null && !root.getMessage().isBlank() - && !message.contains(root.getMessage())) { - message = message + ": " + root.getMessage(); - } - if (isAuthenticationError(e)) { - message = message + ". Hint: authentication failed. Check the RabbitMQ username, " - + "password, and virtual host permissions."; - } - return message; - } - - /** - * Walk the cause chain for an AMQP shutdown signal and map its broker reply - * code to a friendly message. Returns null when there is no AMQP shutdown - * or the reply code has no mapping (caller keeps the raw message). - */ - static String amqpFriendlyMessage(Throwable error) { - for (Throwable current = error; current != null; current = current.getCause()) { - if (!(current instanceof ShutdownSignalException shutdown)) { - continue; - } - Object reason = shutdown.getReason(); - Integer replyCode = null; - String replyText = null; - if (reason instanceof AMQP.Channel.Close channelClose) { - replyCode = channelClose.getReplyCode(); - replyText = channelClose.getReplyText(); - } else if (reason instanceof AMQP.Connection.Close connectionClose) { - replyCode = connectionClose.getReplyCode(); - replyText = connectionClose.getReplyText(); - } - if (replyCode != null) { - String friendly = mapAmqpError(replyCode, replyText); - if (friendly != null) { - return friendly; - } - } - } - return null; - } - - /** Friendly message for a broker reply code, or null to keep the raw message. */ - static String mapAmqpError(int replyCode, String replyText) { - String text = replyText == null ? "" : replyText; - switch (replyCode) { - case 405: { - String name = extractQuotedName(text); - String subject = name != null ? "Queue '" + name + "'" : "The queue"; - return subject + " is exclusive and owned by another connection." - + " Hint: exclusive queues can only be accessed by their owning connection;" - + " stats via the management API are still available."; - } - case 404: { - String name = extractQuotedName(text); - boolean exchange = text.contains("no exchange"); - String kind = exchange ? "Exchange" : "Queue"; - String subject = name != null ? kind + " '" + name + "'" : "The " + kind.toLowerCase(); - return subject + " was not found." - + " Hint: it may have been deleted, or it never existed on this virtual host."; - } - case 406: { - // "PRECONDITION_FAILED - inequivalent arg 'durable' for queue 'q1' - // in vhost '/': ..." — the first quoted token is the argument - // name, so the resource name needs its own extraction. - String name = extractDeclaredResourceName(text); - boolean exchange = text.contains("for exchange"); - String kind = exchange ? "Exchange" : "Queue"; - String subject = name != null ? kind + " '" + name + "'" : "The " + kind.toLowerCase(); - return subject + " already exists with different parameters." - + " Hint: " + kind.toLowerCase() + " parameters are immutable after declaration;" - + " delete and re-declare the " + kind.toLowerCase() + " to change them."; - } - case 403: { - // "ACCESS_REFUSED - access to queue 'q1' in vhost '/' refused for user 'dbx'" - String name = extractQuotedName(text); - String subject = name != null ? "'" + name + "'" : "the requested resource"; - return "Access to " + subject + " was refused." - + " Hint: check the user's configure/write/read permissions on the virtual host."; - } - default: - return null; - } - } - - /** First single-quoted token in a broker reply text (usually the queue/exchange name). */ - static String extractQuotedName(String replyText) { - if (replyText == null) { - return null; - } - Matcher matcher = QUOTED_NAME.matcher(replyText); - return matcher.find() ? matcher.group(1) : null; - } - - /** - * Queue/exchange name in a 406 PRECONDITION_FAILED reply ("... for queue 'q1' - * in vhost ..."); the first quoted token there is the mismatched argument name. - */ - static String extractDeclaredResourceName(String replyText) { - if (replyText == null) { - return null; - } - Matcher matcher = DECLARED_RESOURCE_NAME.matcher(replyText); - return matcher.find() ? matcher.group(2) : null; - } - - private static boolean isAuthenticationError(Throwable error) { - for (Throwable current = error; current != null; current = current.getCause()) { - String className = current.getClass().getName(); - if (className.contains("AuthenticationFailureException") - || className.contains("PossibleAuthenticationFailureException")) { - return true; - } - } - return false; - } - - private static Throwable rootCause(Throwable error) { - Throwable current = error; - for (int depth = 0; current.getCause() != null && current.getCause() != current && depth < 32; depth++) { - current = current.getCause(); - } - return current; - } - - /** - * Channel for the request's effective virtual host. The default vhost reuses - * the primary channel; any other vhost lazily opens (and caches) its own - * connection/channel pair, since AMQP connections are bound to one vhost. - * Channels closed by the broker (e.g. after a 405/404 channel error) are - * detected via {@link #needsNewChannel(Channel)} and rebuilt transparently, - * so one failed call never poisons later ones. - */ - private static Channel channelFor(JsonObject params) throws Exception { - String defaultVhost = cachedConnection != null - ? stringOrDefault(cachedConnection, "virtual_host", "/") : "/"; - String vhost = effectiveVhost(params, cachedConnection); - if (vhost.equals(defaultVhost)) { - return primaryChannel(); - } - - VhostClient client = vhostClients.get(vhost); - if (client != null && client.isOpen()) { - return client.channel; - } - if (client != null) { - client.closeQuietly(); - vhostClients.remove(vhost); - } - JsonObject config = cachedConnection.deepCopy(); - config.addProperty("virtual_host", vhost); - Connection vhostConnection = openConnection(config); - Channel vhostChannel; - try { - vhostChannel = vhostConnection.createChannel(); - } catch (Exception e) { - closeQuietly(vhostConnection); - throw e; - } - vhostClients.put(vhost, new VhostClient(vhostConnection, vhostChannel)); - return vhostChannel; - } - - /** - * Primary channel for the connection's default vhost, recreating the channel - * (or the whole connection) when the broker has closed it. - */ - private static Channel primaryChannel() throws Exception { - if (!needsNewChannel(channel)) { - return channel; - } - if (connection == null || !connection.isOpen()) { - if (cachedConnection == null) { - throw new IllegalStateException("Not connected. Call connect first."); - } - closeQuietly(connection); - connection = openConnection(cachedConnection); - } - closeQuietly(channel); - channel = connection.createChannel(); - return channel; - } - - /** A channel must be rebuilt when it is missing or the broker closed it. */ - static boolean needsNewChannel(Channel ch) { - return ch == null || !ch.isOpen(); - } - - /** - * Effective virtual host: an explicit {@code virtual_host} request parameter - * wins (null/blank means "use the connection's vhost", which is what the - * Rust bridge sends for flat/no-namespace contexts). - */ - static String effectiveVhost(JsonObject params, JsonObject conn) { - String vhost = stringOrNull(params, "virtual_host"); - if (vhost == null || vhost.isBlank()) { - return conn != null ? stringOrDefault(conn, "virtual_host", "/") : "/"; - } - return vhost; - } - - /** - * Whether the request asks for a cross-vhost listing ("all vhosts"). Wins - * over {@code virtual_host}: the vhost-less management API variant is used - * and each returned item carries its own {@code vhost} field. - */ - static boolean allVhostsRequested(JsonObject params) { - return boolOrDefault(params, "all_vhosts", false); - } - - /** - * Management API path for a list endpoint: the vhost-less variant when - * {@code all_vhosts} is set, otherwise scoped to the effective vhost. - */ - static String managementListPath(JsonObject params, JsonObject conn, String resource) { - if (allVhostsRequested(params)) { - return "/api/" + resource; - } - return "/api/" + resource + "/" + urlEncodeVhost(effectiveVhost(params, conn)); - } - - /** - * Client-side vhost filter for connections/channels (the management API - * always lists these cluster-wide); {@code all_vhosts} disables the filter. - * Without an explicit {@code virtual_host} the filter falls back to the - * connection's effective vhost, matching the topic/exchange list behavior. - */ - static String vhostFilter(JsonObject params, JsonObject conn) { - if (allVhostsRequested(params)) { - return ""; - } - return effectiveVhost(params, conn); - } - - /** Copies the source entry's {@code vhost} into the mapped item (all-vhosts listings). */ - static void attachVhost(Map info, JsonObject source) { - info.put("vhost", stringOrEmpty(source, "vhost")); - } - - private static Connection requireConnection() { - if (connection == null) { - throw new IllegalStateException("Not connected. Call connect first."); - } - return connection; - } - - private static JsonObject requireConnectionConfig(JsonObject params) { - JsonObject conn = currentConnectionConfig(params); - if (conn == null) { - throw new IllegalStateException("Not connected. Call connect first."); - } - return conn; - } - - private static JsonObject currentConnectionConfig(JsonObject params) { - if (params.has("connection") && params.get("connection").isJsonObject()) { - return params.getAsJsonObject("connection"); - } - return cachedConnection; - } - - private static JsonObject connectionObject(JsonObject params) { - JsonElement connection = params.get("connection"); - return connection != null && connection.isJsonObject() - ? connection.getAsJsonObject() : params; - } - - /** Queue name: RabbitMQ semantics are flat, so a {@code namespace} parameter is ignored. */ - private static String queueName(JsonObject params) { - String name = stringOrEmpty(params, "topic"); - if (name.isBlank()) { - name = stringOrEmpty(params, "name"); - } - if (name.isBlank()) { - throw new IllegalArgumentException("topic (queue name) is required"); - } - return name; - } - - private static Object argumentValue(JsonElement element) { - if (element == null || !element.isJsonPrimitive()) { - return null; - } - if (element.getAsJsonPrimitive().isBoolean()) { - return element.getAsBoolean(); - } - if (element.getAsJsonPrimitive().isNumber()) { - return element.getAsLong(); - } - return element.getAsString(); - } - - private static String serverString(Map serverProps, String key) { - Object value = serverProps.get(key); - return value != null ? String.valueOf(value) : null; - } - - private static String tryDecodeUtf8(byte[] bytes) { - try { - String text = new String(bytes, StandardCharsets.UTF_8); - // Verify round-trip - byte[] reEncoded = text.getBytes(StandardCharsets.UTF_8); - if (Arrays.equals(bytes, reEncoded)) { - return text; - } - } catch (Exception ignored) {} - return null; - } - - private static String stringOrNull(JsonObject object, String key) { - JsonElement element = object.get(key); - return element == null || element.isJsonNull() ? null : element.getAsString(); - } - - private static String stringOrEmpty(JsonObject object, String key) { - return stringOrDefault(object, key, ""); - } - - private static String stringOrDefault(JsonObject object, String key, String fallback) { - String value = stringOrNull(object, key); - return value == null ? fallback : value; - } - - private static Integer integerOrNull(JsonObject object, String key) { - JsonElement element = object.get(key); - return element == null || element.isJsonNull() ? null : element.getAsInt(); - } - - private static Long longOrNull(JsonObject object, String key) { - JsonElement element = object.get(key); - return element == null || element.isJsonNull() ? null : element.getAsLong(); - } - - private static int intOrDefault(JsonObject object, String key, int fallback) { - Integer value = integerOrNull(object, key); - return value == null ? fallback : value; - } - - private static long longOrDefault(JsonObject object, String key, long fallback) { - Long value = longOrNull(object, key); - return value == null ? fallback : value; - } - - private static boolean boolOrDefault(JsonObject object, String key, boolean fallback) { - JsonElement element = object.get(key); - return element == null || element.isJsonNull() ? fallback : element.getAsBoolean(); - } - - private static Integer integerProperty(JsonObject properties, String key) { - try { - return integerOrNull(properties, key); - } catch (NumberFormatException e) { - return null; - } - } - - private static Boolean booleanProperty(JsonObject properties, String key) { - JsonElement element = properties.get(key); - return element == null || element.isJsonNull() ? null : element.getAsBoolean(); - } - - private static boolean boolProperty(JsonObject conn, String key) { - JsonObject properties = conn.has("properties") && conn.get("properties").isJsonObject() - ? conn.getAsJsonObject("properties") : null; - return properties != null && boolOrDefault(properties, key, false); - } - - // ----------------------------------------------------------------------- - // Inner types - // ----------------------------------------------------------------------- - - private static final class HandshakeResult { - private final int protocolVersion; - private final int agentProtocolVersion; - private final List capabilities; - - private HandshakeResult(int protocolVersion, int agentProtocolVersion, List capabilities) { - this.protocolVersion = protocolVersion; - this.agentProtocolVersion = agentProtocolVersion; - this.capabilities = capabilities; - } - } - - /** Connection/channel pair cached for one non-default virtual host. */ - private static final class VhostClient { - private final Connection connection; - private final Channel channel; - - private VhostClient(Connection connection, Channel channel) { - this.connection = connection; - this.channel = channel; - } - - private boolean isOpen() { - return connection.isOpen() && channel.isOpen(); - } - - private void closeQuietly() { - RabbitMqAgent.closeQuietly(channel); - RabbitMqAgent.closeQuietly(connection); - } - } -} diff --git a/agents/drivers/rabbitmq/src/test/java/com/dbx/agent/rabbitmq/RabbitMqAgentTest.java b/agents/drivers/rabbitmq/src/test/java/com/dbx/agent/rabbitmq/RabbitMqAgentTest.java deleted file mode 100644 index cdb285046..000000000 --- a/agents/drivers/rabbitmq/src/test/java/com/dbx/agent/rabbitmq/RabbitMqAgentTest.java +++ /dev/null @@ -1,1779 +0,0 @@ -package com.dbx.agent.rabbitmq; - -import static org.junit.jupiter.api.Assertions.assertEquals; -import static org.junit.jupiter.api.Assertions.assertFalse; -import static org.junit.jupiter.api.Assertions.assertNull; -import static org.junit.jupiter.api.Assertions.assertThrows; -import static org.junit.jupiter.api.Assertions.assertTrue; - -import com.google.gson.JsonArray; -import com.google.gson.JsonObject; -import com.google.gson.JsonParser; -import com.rabbitmq.client.AMQP; -import com.rabbitmq.client.Address; -import com.rabbitmq.client.Channel; -import com.rabbitmq.client.ConnectionFactory; -import com.rabbitmq.client.ShutdownSignalException; -import com.rabbitmq.client.impl.AMQImpl; -import com.sun.net.httpserver.HttpServer; -import java.lang.reflect.Proxy; -import java.net.InetSocketAddress; -import java.nio.charset.StandardCharsets; -import java.util.ArrayList; -import java.util.List; -import java.util.Map; -import org.junit.jupiter.api.Test; - -class RabbitMqAgentTest { - - // ------------------------------------------------------------------- - // Address parsing - // ------------------------------------------------------------------- - - @Test - void parsesCommaSeparatedHostPortPairs() { - List
addresses = RabbitMqAgent.parseAddresses("a:5672, b:5673", 5672); - assertEquals(2, addresses.size()); - assertEquals("a", addresses.get(0).getHost()); - assertEquals(5672, addresses.get(0).getPort()); - assertEquals("b", addresses.get(1).getHost()); - assertEquals(5673, addresses.get(1).getPort()); - } - - @Test - void bareHostFallsBackToDefaultPort() { - List
addresses = RabbitMqAgent.parseAddresses("rabbit.internal", 5672); - assertEquals(1, addresses.size()); - assertEquals("rabbit.internal", addresses.get(0).getHost()); - assertEquals(5672, addresses.get(0).getPort()); - } - - @Test - void skipsBlankAddressEntries() { - List
addresses = RabbitMqAgent.parseAddresses(" a:5672,, ", 5672); - assertEquals(1, addresses.size()); - assertEquals("a", addresses.get(0).getHost()); - } - - @Test - void rejectsBlankAddressList() { - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.parseAddresses(" , ", 5672)); - } - - @Test - void resolveAddressesUsesPortParameterForBareHosts() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "host1,host2:5673", "port": 5670 } - """).getAsJsonObject(); - List
addresses = RabbitMqAgent.resolveAddresses(conn); - assertEquals(2, addresses.size()); - assertEquals(5670, addresses.get(0).getPort()); - assertEquals(5673, addresses.get(1).getPort()); - } - - @Test - void resolveAddressesDefaultsToAmqpPort() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "host1" } - """).getAsJsonObject(); - List
addresses = RabbitMqAgent.resolveAddresses(conn); - assertEquals(5672, addresses.get(0).getPort()); - } - - @Test - void resolveAddressesRequiresAddresses() { - JsonObject conn = JsonParser.parseString(""" - { "username": "guest" } - """).getAsJsonObject(); - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.resolveAddresses(conn)); - } - - // ------------------------------------------------------------------- - // Peek normalization - // ------------------------------------------------------------------- - - @Test - void normalizesNegativePeekOffsetToZero() { - assertEquals(0L, RabbitMqAgent.normalizePeekOffset(-5)); - } - - @Test - void keepsPositivePeekOffset() { - assertEquals(7L, RabbitMqAgent.normalizePeekOffset(7)); - } - - @Test - void normalizesPeekCountToAtLeastOne() { - assertEquals(1, RabbitMqAgent.normalizePeekCount(0)); - assertEquals(1, RabbitMqAgent.normalizePeekCount(-3)); - assertEquals(10, RabbitMqAgent.normalizePeekCount(10)); - } - - // ------------------------------------------------------------------- - // Send routing key resolution - // ------------------------------------------------------------------- - - @Test - void routingKeyFallsBackToMessageKeyThenQueue() { - assertEquals("q1", RabbitMqAgent.resolveRoutingKey(JsonParser.parseString(""" - { "topic": "q1" } - """).getAsJsonObject(), "q1")); - assertEquals("orders.new", RabbitMqAgent.resolveRoutingKey(JsonParser.parseString(""" - { "topic": "q1", "key": "orders.new" } - """).getAsJsonObject(), "q1")); - assertEquals("rk", RabbitMqAgent.resolveRoutingKey(JsonParser.parseString(""" - { "topic": "q1", "key": "orders.new", "routing_key": "rk" } - """).getAsJsonObject(), "q1")); - assertEquals("rk2", RabbitMqAgent.resolveRoutingKey(JsonParser.parseString(""" - { "topic": "q1", "routingKey": "rk2" } - """).getAsJsonObject(), "q1")); - } - - @Test - void blankRoutingKeyFallsBackToQueue() { - assertEquals("q1", RabbitMqAgent.resolveRoutingKey(JsonParser.parseString(""" - { "topic": "q1", "key": "" } - """).getAsJsonObject(), "q1")); - assertEquals("q1", RabbitMqAgent.resolveRoutingKey(JsonParser.parseString(""" - { "topic": "q1", "key": " " } - """).getAsJsonObject(), "q1")); - assertEquals("q1", RabbitMqAgent.resolveRoutingKey(JsonParser.parseString(""" - { "topic": "q1", "routing_key": "", "key": "" } - """).getAsJsonObject(), "q1")); - } - - // ------------------------------------------------------------------- - // Connection factory - // ------------------------------------------------------------------- - - @Test - void buildsConnectionFactoryWithDefaults() throws Exception { - ConnectionFactory factory = RabbitMqAgent.buildConnectionFactory(JsonParser.parseString(""" - { "addresses": "localhost" } - """).getAsJsonObject()); - assertEquals("guest", factory.getUsername()); - assertEquals("/", factory.getVirtualHost()); - } - - @Test - void buildsConnectionFactoryWithVirtualHostAndCredentials() throws Exception { - ConnectionFactory factory = RabbitMqAgent.buildConnectionFactory(JsonParser.parseString(""" - { "addresses": "localhost", "username": "dbx", "password": "secret", "virtual_host": "/tenant" } - """).getAsJsonObject()); - assertEquals("dbx", factory.getUsername()); - assertEquals("/tenant", factory.getVirtualHost()); - } - - @Test - void appliesExtraPropertiesToConnectionFactory() throws Exception { - ConnectionFactory factory = RabbitMqAgent.buildConnectionFactory(JsonParser.parseString(""" - { - "addresses": "localhost", - "properties": { - "requested_heartbeat": 30, - "connection_timeout_ms": 5000, - "automatic_recovery": false - } - } - """).getAsJsonObject()); - assertEquals(30, factory.getRequestedHeartbeat()); - assertEquals(5000, factory.getConnectionTimeout()); - assertFalse(factory.isAutomaticRecoveryEnabled()); - } - - @Test - void enablesTlsWithoutVerificationWhenSkipVerifyRequested() throws Exception { - ConnectionFactory factory = RabbitMqAgent.buildConnectionFactory(JsonParser.parseString(""" - { "addresses": "localhost", "tls_skip_verify": true } - """).getAsJsonObject()); - assertTrue(factory.isSSL()); - } - - // ------------------------------------------------------------------- - // Management API helpers - // ------------------------------------------------------------------- - - @Test - void buildsBasicAuthHeader() { - assertEquals("Basic Z3Vlc3Q6Z3Vlc3Q=", RabbitMqAgent.basicAuthHeader("guest", "guest")); - } - - @Test - void buildsManagementBaseUrl() { - assertEquals("http://localhost:15672", RabbitMqAgent.managementBaseUrl("localhost", 15672, false)); - assertEquals("https://mq:15671", RabbitMqAgent.managementBaseUrl("mq", 15671, true)); - } - - @Test - void managementPortDefaultsTo15672Or15671ForTls() { - JsonObject plain = JsonParser.parseString(""" - { "addresses": "localhost" } - """).getAsJsonObject(); - assertEquals(15672, RabbitMqAgent.managementPort(plain, false)); - assertEquals(15671, RabbitMqAgent.managementPort(plain, true)); - } - - @Test - void managementPortCanBeOverriddenViaProperties() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost", "properties": { "management_port": 55672 } } - """).getAsJsonObject(); - assertEquals(55672, RabbitMqAgent.managementPort(conn, false)); - } - - @Test - void managementErrorMessageBlamesCredentialsOn401And403() { - for (int status : new int[] {401, 403}) { - String message = RabbitMqAgent.managementErrorMessage(status, "GET", "/api/queues"); - assertTrue(message.contains("HTTP " + status), message); - assertTrue(message.contains("management permission tag"), message); - assertFalse(message.contains("rabbitmq_management plugin must be enabled"), message); - } - } - - @Test - void managementErrorMessageKeepsPluginHintForOtherStatuses() { - String message = RabbitMqAgent.managementErrorMessage(404, "GET", "/api/queues/%2F/gone"); - assertTrue(message.contains("HTTP 404")); - assertTrue(message.contains("rabbitmq_management plugin must be enabled")); - assertFalse(message.contains("management permission tag")); - } - - @Test - void managementRequestSurfaces401AsCredentialError() throws Exception { - // A local stub server standing in for the management API: the agent must - // attribute a 401 to credentials/permissions, not to a missing plugin. - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/api", exchange -> { - exchange.sendResponseHeaders(401, -1); - exchange.close(); - }); - server.start(); - try { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "127.0.0.1", "properties": { "management_port": %d } } - """.formatted(server.getAddress().getPort())).getAsJsonObject(); - Exception error = assertThrows(IllegalStateException.class, - () -> RabbitMqAgent.managementGet(conn, "/api/queues")); - assertTrue(error.getMessage().contains("HTTP 401"), error.getMessage()); - assertTrue(error.getMessage().contains("management permission tag"), error.getMessage()); - assertFalse(error.getMessage().contains("plugin must be enabled"), error.getMessage()); - } finally { - server.stop(0); - } - } - - @Test - void explicitManagementUrlIsUsedVerbatimWithTrailingSlashTrimmed() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "mq1:5672,mq2:5672", "management_url": "https://proxy:8443/rmq/" } - """).getAsJsonObject(); - assertEquals(List.of("https://proxy:8443/rmq"), RabbitMqAgent.managementBaseUrls(conn)); - } - - @Test - void explicitManagementUrlDoesNotRequireAddresses() { - JsonObject conn = JsonParser.parseString(""" - { "management_url": "http://mgmt:15672" } - """).getAsJsonObject(); - assertEquals(List.of("http://mgmt:15672"), RabbitMqAgent.managementBaseUrls(conn)); - } - - @Test - void derivedManagementBaseUrlsCoverAllAddresses() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "mq1:5672,mq2:5673" } - """).getAsJsonObject(); - assertEquals(List.of("http://mq1:15672", "http://mq2:15672"), - RabbitMqAgent.managementBaseUrls(conn)); - } - - @Test - void tlsSkipVerifyAloneDoesNotFlipDerivedSchemeToHttps() { - // tls_skip_verify is a verification flag, not a scheme indicator. - JsonObject skipVerifyOnly = JsonParser.parseString(""" - { "addresses": "mq1", "tls_skip_verify": true } - """).getAsJsonObject(); - assertEquals(List.of("http://mq1:15672"), RabbitMqAgent.managementBaseUrls(skipVerifyOnly)); - - JsonObject tlsObject = JsonParser.parseString(""" - { "addresses": "mq1", "tls": { "skip_verify": true } } - """).getAsJsonObject(); - assertEquals(List.of("https://mq1:15671"), RabbitMqAgent.managementBaseUrls(tlsObject)); - - JsonObject sslProperty = JsonParser.parseString(""" - { "addresses": "mq1", "properties": { "ssl": true } } - """).getAsJsonObject(); - assertEquals(List.of("https://mq1:15671"), RabbitMqAgent.managementBaseUrls(sslProperty)); - } - - @Test - void blankCredentialsFallBackToGuest() throws Exception { - ConnectionFactory factory = RabbitMqAgent.buildConnectionFactory(JsonParser.parseString(""" - { "addresses": "localhost", "username": "", "password": " " } - """).getAsJsonObject()); - assertEquals("guest", factory.getUsername()); - assertEquals("guest", factory.getPassword()); - - ConnectionFactory nullCredentials = RabbitMqAgent.buildConnectionFactory(JsonParser.parseString(""" - { "addresses": "localhost", "username": null } - """).getAsJsonObject()); - assertEquals("guest", nullCredentials.getUsername()); - } - - @Test - void managementGetAllPaginatesToLastPage() throws Exception { - List requestedPages = new ArrayList<>(); - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/api/queues", exchange -> { - String query = exchange.getRequestURI().getQuery(); - int page = Integer.parseInt(query.replaceAll(".*page=(\\d+).*", "$1")); - requestedPages.add(page); - byte[] body = switch (page) { - case 1 -> """ - { "items": [ { "name": "q1" } ], "page": 1, "page_count": 3, "total_count": 3 } - """.getBytes(StandardCharsets.UTF_8); - case 2 -> """ - { "items": [ { "name": "q2" } ], "page": 2, "page_count": 3, "total_count": 3 } - """.getBytes(StandardCharsets.UTF_8); - default -> """ - { "items": [ { "name": "q3" } ], "page": 3, "page_count": 3, "total_count": 3 } - """.getBytes(StandardCharsets.UTF_8); - }; - exchange.getResponseHeaders().add("Content-Type", "application/json"); - exchange.sendResponseHeaders(200, body.length); - exchange.getResponseBody().write(body); - exchange.close(); - }); - server.start(); - try { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "127.0.0.1", "properties": { "management_port": %d } } - """.formatted(server.getAddress().getPort())).getAsJsonObject(); - JsonArray all = RabbitMqAgent.managementGetAll(conn, "/api/queues"); - assertEquals(3, all.size()); - assertEquals("q1", all.get(0).getAsJsonObject().get("name").getAsString()); - assertEquals("q3", all.get(2).getAsJsonObject().get("name").getAsString()); - assertEquals(List.of(1, 2, 3), requestedPages); - } finally { - server.stop(0); - } - } - - @Test - void managementGetAllAcceptsPlainArrayResponse() throws Exception { - List requests = new ArrayList<>(); - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/api/users", exchange -> { - requests.add(exchange.getRequestURI().toString()); - byte[] body = """ - [ { "name": "guest", "tags": "administrator" } ] - """.getBytes(StandardCharsets.UTF_8); - exchange.getResponseHeaders().add("Content-Type", "application/json"); - exchange.sendResponseHeaders(200, body.length); - exchange.getResponseBody().write(body); - exchange.close(); - }); - server.start(); - try { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "127.0.0.1", "properties": { "management_port": %d } } - """.formatted(server.getAddress().getPort())).getAsJsonObject(); - JsonArray all = RabbitMqAgent.managementGetAll(conn, "/api/users"); - assertEquals(1, all.size()); - // A plain-array answer means the broker ignored pagination: stop there. - assertEquals(1, requests.size()); - } finally { - server.stop(0); - } - } - - @Test - void managementRequestFailsOverAcrossDerivedCandidates() throws Exception { - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/api/queues", exchange -> { - byte[] body = "[]".getBytes(StandardCharsets.UTF_8); - exchange.getResponseHeaders().add("Content-Type", "application/json"); - exchange.sendResponseHeaders(200, body.length); - exchange.getResponseBody().write(body); - exchange.close(); - }); - server.start(); - try { - // 127.0.0.2 refuses the connection; the second candidate answers. - JsonObject conn = JsonParser.parseString(""" - { "addresses": "127.0.0.2,127.0.0.1", "properties": { "management_port": %d } } - """.formatted(server.getAddress().getPort())).getAsJsonObject(); - assertTrue(RabbitMqAgent.managementGet(conn, "/api/queues").isJsonArray()); - } finally { - server.stop(0); - } - } - - @Test - void httpErrorStatusDoesNotTriggerFailover() throws Exception { - HttpServer rejecting = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - rejecting.createContext("/api", exchange -> { - exchange.sendResponseHeaders(401, -1); - exchange.close(); - }); - rejecting.start(); - int port = rejecting.getAddress().getPort(); - try { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "127.0.0.1,127.0.0.2", "properties": { "management_port": %d } } - """.formatted(port)).getAsJsonObject(); - Exception error = assertThrows(IllegalStateException.class, - () -> RabbitMqAgent.managementGet(conn, "/api/queues")); - // If the second candidate were attempted, its connection failure - // would replace this terminal HTTP status with an I/O error. - assertTrue(error.getMessage().contains("HTTP 401"), error.getMessage()); - } finally { - rejecting.stop(0); - } - } - - @Test - void managementUrlWithPathPrefixReachesStub() throws Exception { - List requests = new ArrayList<>(); - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/rmq/api/queues", exchange -> { - requests.add(exchange.getRequestURI().getRawPath()); - byte[] body = """ - [ { "name": "dbx-q1", "durable": true, "state": "running" } ] - """.getBytes(StandardCharsets.UTF_8); - exchange.getResponseHeaders().add("Content-Type", "application/json"); - exchange.sendResponseHeaders(200, body.length); - exchange.getResponseBody().write(body); - exchange.close(); - }); - server.start(); - try { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 70, "method": "mq_list_topics", - "params": { "connection": { "addresses": "192.0.2.1:5672", - "management_url": "http://127.0.0.1:%d/rmq/" } } } - """.formatted(server.getAddress().getPort()))).getAsJsonObject(); - JsonArray topics = response.getAsJsonObject("result").getAsJsonArray("topics"); - assertEquals("dbx-q1", topics.get(0).getAsJsonObject().get("name").getAsString()); - // The reverse-proxy path prefix is preserved verbatim. - assertEquals("/rmq/api/queues/%2F", requests.get(0)); - } finally { - server.stop(0); - } - } - - @Test - void encodesDefaultVhostForManagementApi() { - assertEquals("%2F", RabbitMqAgent.urlEncodeVhost("/")); - assertEquals("tenant-a", RabbitMqAgent.urlEncodeVhost("tenant-a")); - } - - @Test - void urlEncodePathSegmentEncodesSpacesAsPercent20() { - // URLEncoder's form-style '+' for spaces 404s on the management API. - assertEquals("dbx-space%20test", RabbitMqAgent.urlEncodePathSegment("dbx-space test")); - assertEquals("plain-name", RabbitMqAgent.urlEncodePathSegment("plain-name")); - assertEquals("a%2Fb%3Ac", RabbitMqAgent.urlEncodePathSegment("a/b:c")); - assertEquals("%E4%B8%AD%E6%96%87%20queue", RabbitMqAgent.urlEncodePathSegment("中文 queue")); - } - - @Test - void tlsSkipVerifyReadsTopLevelAndNestedFlags() { - assertFalse(RabbitMqAgent.tlsSkipVerify(JsonParser.parseString(""" - { "addresses": "localhost" } - """).getAsJsonObject())); - assertTrue(RabbitMqAgent.tlsSkipVerify(JsonParser.parseString(""" - { "addresses": "localhost", "tls_skip_verify": true } - """).getAsJsonObject())); - assertTrue(RabbitMqAgent.tlsSkipVerify(JsonParser.parseString(""" - { "addresses": "localhost", "tls": { "skip_verify": true } } - """).getAsJsonObject())); - assertFalse(RabbitMqAgent.tlsSkipVerify(JsonParser.parseString(""" - { "addresses": "localhost", "tls": { "skip_verify": false } } - """).getAsJsonObject())); - } - - // ------------------------------------------------------------------- - // JSON-RPC envelope (no broker required) - // ------------------------------------------------------------------- - - @Test - void handshakeReportsCapabilities() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 1, "method": "handshake", "params": {} } - """); - JsonObject result = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("result"); - assertEquals(1, result.get("protocolVersion").getAsInt()); - assertTrue(result.getAsJsonArray("capabilities").toString().contains("mq_topics")); - assertTrue(result.getAsJsonArray("capabilities").toString().contains("mq_messages")); - } - - @Test - void unknownMethodReturnsError() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 2, "method": "mq_bogus", "params": {} } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Unknown method")); - } - - @Test - void topicOperationsRequireConnection() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 3, "method": "mq_create_topic", "params": { "name": "q1" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Not connected")); - } - - @Test - void alterTopicConfigIsRejectedAsUnsupported() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 4, "method": "mq_alter_topic_config", "params": { "name": "q1", "configs": [] } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("immutable")); - } - - @Test - void testConnectionFailsFastWithoutAddresses() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 5, "method": "test_connection", "params": { "connection": { "addresses": "" } } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("addresses is required")); - } - - @Test - void malformedRequestReturnsErrorWithNullId() { - // A request that is not valid JSON must not kill the agent process. - JsonObject response = JsonParser.parseString( - RabbitMqAgent.handleRequest("this is not json")).getAsJsonObject(); - assertTrue(response.get("id").isJsonNull()); - assertEquals(-1, response.getAsJsonObject("error").get("code").getAsInt()); - - JsonObject notAnObject = JsonParser.parseString( - RabbitMqAgent.handleRequest("[1, 2, 3]")).getAsJsonObject(); - assertTrue(notAnObject.get("id").isJsonNull()); - assertTrue(notAnObject.has("error")); - } - - @Test - void missingMethodReturnsErrorButKeepsRequestId() { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 6, "params": {} } - """)).getAsJsonObject(); - assertEquals(6, response.get("id").getAsInt()); - assertEquals(-1, response.getAsJsonObject("error").get("code").getAsInt()); - } - - // ------------------------------------------------------------------- - // all_vhosts fail-fast on non-list operations - // ------------------------------------------------------------------- - - @Test - void allVhostsIsRejectedForNonListOperations() { - List methods = List.of( - "mq_create_topic", "mq_delete_topic", "mq_purge_queue", "mq_send_message", - "mq_bind", "mq_unbind", "mq_create_exchange", "mq_delete_exchange", - "mq_peek_messages", "mq_get_topic_stats", "mq_list_consumers", "mq_close_connection"); - for (String method : methods) { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 40, "method": "%s", - "params": { "all_vhosts": true, "topic": "q1", "name": "q1", - "source": "ex1", "destination": "q1", "destinationType": "queue", - "type": "direct" } } - """.formatted(method))).getAsJsonObject(); - JsonObject error = response.getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt(), method); - assertEquals("all_vhosts is only supported for list operations", - error.get("message").getAsString(), method); - } - } - - @Test - void allVhostsRejectionPrecedesConnectionCheck() { - // The semantic error must win over "Not connected": no broker is needed. - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 41, "method": "mq_purge_queue", - "params": { "all_vhosts": true, "topic": "q1" } } - """); - String message = JsonParser.parseString(response).getAsJsonObject() - .getAsJsonObject("error").get("message").getAsString(); - assertTrue(message.contains("all_vhosts is only supported for list operations")); - assertFalse(message.contains("Not connected")); - } - - @Test - void listOperationsStillAcceptAllVhosts() { - // Without a connection these fail with "Not connected", proving the - // all_vhosts guard did not reject them first. - for (String method : List.of("mq_list_topics", "mq_list_exchanges", "mq_list_bindings", - "mq_list_connections", "mq_list_channels")) { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 42, "method": "%s", "params": { "all_vhosts": true } } - """.formatted(method))).getAsJsonObject(); - String message = response.getAsJsonObject("error").get("message").getAsString(); - assertTrue(message.contains("Not connected"), method + ": " + message); - } - } - - // ------------------------------------------------------------------- - // Effective virtual host resolution - // ------------------------------------------------------------------- - - @Test - void explicitVirtualHostParameterWins() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost", "virtual_host": "/default" } - """).getAsJsonObject(); - JsonObject params = JsonParser.parseString(""" - { "topic": "q1", "virtual_host": "/tenant" } - """).getAsJsonObject(); - assertEquals("/tenant", RabbitMqAgent.effectiveVhost(params, conn)); - } - - @Test - void blankVirtualHostFallsBackToConnectionVhost() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost", "virtual_host": "/default" } - """).getAsJsonObject(); - assertEquals("/default", RabbitMqAgent.effectiveVhost(JsonParser.parseString(""" - { "topic": "q1", "virtual_host": "" } - """).getAsJsonObject(), conn)); - assertEquals("/default", RabbitMqAgent.effectiveVhost(JsonParser.parseString(""" - { "topic": "q1", "virtual_host": " " } - """).getAsJsonObject(), conn)); - assertEquals("/default", RabbitMqAgent.effectiveVhost(JsonParser.parseString(""" - { "topic": "q1", "virtual_host": null } - """).getAsJsonObject(), conn)); - } - - @Test - void missingVirtualHostFallsBackToConnectionThenSlash() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost", "virtual_host": "/default" } - """).getAsJsonObject(); - JsonObject params = JsonParser.parseString(""" - { "topic": "q1" } - """).getAsJsonObject(); - assertEquals("/default", RabbitMqAgent.effectiveVhost(params, conn)); - assertEquals("/", RabbitMqAgent.effectiveVhost(params, null)); - JsonObject noVhostConn = JsonParser.parseString(""" - { "addresses": "localhost" } - """).getAsJsonObject(); - assertEquals("/", RabbitMqAgent.effectiveVhost(params, noVhostConn)); - } - - // ------------------------------------------------------------------- - // All-vhosts listing - // ------------------------------------------------------------------- - - @Test - void allVhostsRequestedDefaultsToFalse() { - assertFalse(RabbitMqAgent.allVhostsRequested(JsonParser.parseString(""" - { "virtual_host": "/tenant" } - """).getAsJsonObject())); - assertFalse(RabbitMqAgent.allVhostsRequested(JsonParser.parseString(""" - { "all_vhosts": false } - """).getAsJsonObject())); - assertTrue(RabbitMqAgent.allVhostsRequested(JsonParser.parseString(""" - { "all_vhosts": true } - """).getAsJsonObject())); - } - - @Test - void managementListPathScopesToEffectiveVhostByDefault() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost", "virtual_host": "/default" } - """).getAsJsonObject(); - assertEquals("/api/queues/%2Fdefault", RabbitMqAgent.managementListPath( - JsonParser.parseString("{}").getAsJsonObject(), conn, "queues")); - assertEquals("/api/exchanges/%2Ftenant", RabbitMqAgent.managementListPath( - JsonParser.parseString(""" - { "virtual_host": "/tenant" } - """).getAsJsonObject(), conn, "exchanges")); - JsonObject noVhostConn = JsonParser.parseString(""" - { "addresses": "localhost" } - """).getAsJsonObject(); - assertEquals("/api/bindings/%2F", RabbitMqAgent.managementListPath( - JsonParser.parseString("{}").getAsJsonObject(), noVhostConn, "bindings")); - } - - @Test - void managementListPathUsesVhostlessVariantWhenAllVhosts() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost", "virtual_host": "/default" } - """).getAsJsonObject(); - JsonObject params = JsonParser.parseString(""" - { "all_vhosts": true } - """).getAsJsonObject(); - assertEquals("/api/queues", RabbitMqAgent.managementListPath(params, conn, "queues")); - assertEquals("/api/exchanges", RabbitMqAgent.managementListPath(params, conn, "exchanges")); - assertEquals("/api/bindings", RabbitMqAgent.managementListPath(params, conn, "bindings")); - } - - @Test - void allVhostsWinsOverExplicitVirtualHost() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost" } - """).getAsJsonObject(); - JsonObject params = JsonParser.parseString(""" - { "all_vhosts": true, "virtual_host": "/tenant" } - """).getAsJsonObject(); - assertEquals("/api/queues", RabbitMqAgent.managementListPath(params, conn, "queues")); - assertEquals("", RabbitMqAgent.vhostFilter(params, conn)); - } - - @Test - void vhostFilterPassesThroughExplicitVirtualHost() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost", "virtual_host": "/default" } - """).getAsJsonObject(); - assertEquals("/tenant", RabbitMqAgent.vhostFilter(JsonParser.parseString(""" - { "virtual_host": "/tenant" } - """).getAsJsonObject(), conn)); - } - - @Test - void vhostFilterFallsBackToConnectionVhost() { - JsonObject conn = JsonParser.parseString(""" - { "addresses": "localhost", "virtual_host": "/default" } - """).getAsJsonObject(); - JsonObject params = JsonParser.parseString("{}").getAsJsonObject(); - assertEquals("/default", RabbitMqAgent.vhostFilter(params, conn)); - assertEquals("/", RabbitMqAgent.vhostFilter(params, null)); - JsonObject noVhostConn = JsonParser.parseString(""" - { "addresses": "localhost" } - """).getAsJsonObject(); - assertEquals("/", RabbitMqAgent.vhostFilter(params, noVhostConn)); - } - - @Test - void attachVhostCopiesSourceVhost() { - java.util.Map info = new java.util.LinkedHashMap<>(); - RabbitMqAgent.attachVhost(info, JsonParser.parseString(""" - { "name": "q1", "vhost": "/tenant-a" } - """).getAsJsonObject()); - assertEquals("/tenant-a", info.get("vhost")); - - java.util.Map missing = new java.util.LinkedHashMap<>(); - RabbitMqAgent.attachVhost(missing, JsonParser.parseString(""" - { "name": "q1" } - """).getAsJsonObject()); - assertEquals("", missing.get("vhost")); - } - - // ------------------------------------------------------------------- - // consumer_details mapping - // ------------------------------------------------------------------- - - @Test - void mapsConsumerDetailsFromQueueInfo() { - JsonObject info = JsonParser.parseString(""" - { - "name": "q1", - "consumer_details": [ - { - "consumer_tag": "amq.ctag-abc", - "ack_required": true, - "prefetch_count": 20, - "active": true, - "channel_details": { "name": "10.0.0.1:5672 -> 10.0.0.2:41234 (1)", "number": 1 } - }, - { - "consumer_tag": "amq.ctag-def", - "ack_required": false, - "active": false - } - ] - } - """).getAsJsonObject(); - var consumers = RabbitMqAgent.consumersFromQueueInfo(info); - assertEquals(2, consumers.size()); - - var first = consumers.get(0); - assertEquals("10.0.0.1:5672 -> 10.0.0.2:41234 (1)", first.get("name")); - assertEquals("amq.ctag-abc", first.get("tag")); - assertEquals(true, first.get("active")); - assertEquals(true, first.get("ackRequired")); - assertEquals(20, first.get("prefetch")); - - var second = consumers.get(1); - assertEquals("", second.get("name")); - assertEquals("amq.ctag-def", second.get("tag")); - assertEquals(false, second.get("active")); - assertEquals(false, second.get("ackRequired")); - assertFalse(second.containsKey("prefetch")); - } - - @Test - void missingConsumerDetailsMapsToEmptyList() { - assertTrue(RabbitMqAgent.consumersFromQueueInfo(JsonParser.parseString(""" - { "name": "q1" } - """).getAsJsonObject()).isEmpty()); - assertTrue(RabbitMqAgent.consumersFromQueueInfo(JsonParser.parseString(""" - { "name": "q1", "consumer_details": [] } - """).getAsJsonObject()).isEmpty()); - } - - // ------------------------------------------------------------------- - // Purge queue - // ------------------------------------------------------------------- - - @Test - void purgeQueueRequiresTopic() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 10, "method": "mq_purge_queue", "params": {} } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("topic (queue name) is required")); - } - - @Test - void purgeQueueRequiresConnection() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 11, "method": "mq_purge_queue", "params": { "topic": "q1" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Not connected")); - } - - // ------------------------------------------------------------------- - // Consumers / namespaces (no broker required) - // ------------------------------------------------------------------- - - @Test - void listConsumersRequiresConnection() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 12, "method": "mq_list_consumers", "params": { "topic": "q1" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Not connected")); - } - - @Test - void listNamespacesRequiresConnection() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 13, "method": "mq_list_namespaces", "params": {} } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Not connected")); - } - - @Test - void createNamespaceRequiresName() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 14, "method": "mq_create_namespace", "params": { "namespace": " " } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("namespace is required")); - } - - @Test - void createNamespaceRejectsAllVhostsMarker() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 14, "method": "mq_create_namespace", "params": { "namespace": "*" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("all-vhosts")); - } - - @Test - void deleteNamespaceRejectsAllVhostsMarker() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 15, "method": "mq_delete_namespace", "params": { "namespace": "*" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("all-vhosts")); - } - - @Test - void deleteNamespaceRejectsDefaultVhost() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 15, "method": "mq_delete_namespace", "params": { "namespace": "/" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("cannot be deleted")); - } - - @Test - void deleteNamespaceGuardRejectsConnectedVhost() { - assertThrows(IllegalArgumentException.class, - () -> RabbitMqAgent.assertNamespaceDeletable("/", null)); - assertThrows(IllegalArgumentException.class, - () -> RabbitMqAgent.assertNamespaceDeletable("/tenant", "/tenant")); - RabbitMqAgent.assertNamespaceDeletable("/tenant", "/"); - RabbitMqAgent.assertNamespaceDeletable("dbx-tier1-vhost", null); - } - - // ------------------------------------------------------------------- - // Exchanges & bindings - // ------------------------------------------------------------------- - - @Test - void validatesExchangeTypeWhitelist() { - assertEquals("direct", RabbitMqAgent.validateExchangeType("direct")); - assertEquals("fanout", RabbitMqAgent.validateExchangeType("fanout")); - assertEquals("topic", RabbitMqAgent.validateExchangeType("topic")); - assertEquals("headers", RabbitMqAgent.validateExchangeType("headers")); - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.validateExchangeType("")); - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.validateExchangeType("x-delayed-message")); - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.validateExchangeType("Direct")); - } - - @Test - void exchangeDeletionGuardRejectsDefaultAndBuiltIns() { - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.assertExchangeDeletable("")); - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.assertExchangeDeletable("amq.direct")); - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.assertExchangeDeletable("amq.topic")); - RabbitMqAgent.assertExchangeDeletable("dbx-ex-change"); - RabbitMqAgent.assertExchangeDeletable("amqp.custom"); - } - - @Test - void mapsExchangeInfoWithDefaultExchangeType() { - var defaultExchange = RabbitMqAgent.exchangeInfoFromJson(JsonParser.parseString(""" - { "name": "", "type": "", "durable": true, "auto_delete": false, "internal": false } - """).getAsJsonObject()); - assertEquals("", defaultExchange.get("name")); - assertEquals("default", defaultExchange.get("type")); - assertEquals(true, defaultExchange.get("durable")); - assertEquals(false, defaultExchange.get("autoDelete")); - assertEquals(false, defaultExchange.get("internal")); - } - - @Test - void mapsExchangeInfoKeepsDeclaredType() { - var exchange = RabbitMqAgent.exchangeInfoFromJson(JsonParser.parseString(""" - { "name": "amq.topic", "type": "topic", "durable": true, "auto_delete": false, "internal": false } - """).getAsJsonObject()); - assertEquals("amq.topic", exchange.get("name")); - assertEquals("topic", exchange.get("type")); - } - - @Test - void mapsBindingInfoToCamelCase() { - var binding = RabbitMqAgent.bindingInfoFromJson(JsonParser.parseString(""" - { - "source": "dbx-ex-change", - "destination": "dbx-ex-test", - "destination_type": "queue", - "routing_key": "dbx.key", - "arguments": { "x-match": "all", "retries": 3, "drop": null } - } - """).getAsJsonObject()); - assertEquals("dbx-ex-change", binding.get("source")); - assertEquals("dbx-ex-test", binding.get("destination")); - assertEquals("queue", binding.get("destinationType")); - assertEquals("dbx.key", binding.get("routingKey")); - var arguments = (java.util.Map) binding.get("arguments"); - assertEquals("all", arguments.get("x-match")); - assertEquals(3L, arguments.get("retries")); - assertFalse(arguments.containsKey("drop")); - } - - @Test - void bindingInfoOmitsEmptyArguments() { - var binding = RabbitMqAgent.bindingInfoFromJson(JsonParser.parseString(""" - { - "source": "ex1", - "destination": "ex2", - "destination_type": "exchange", - "routing_key": "", - "arguments": {} - } - """).getAsJsonObject()); - assertEquals("exchange", binding.get("destinationType")); - assertFalse(binding.containsKey("arguments")); - } - - @Test - void createExchangeRejectsInvalidTypeBeforeConnecting() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 20, "method": "mq_create_exchange", - "params": { "name": "ex1", "type": "bogus" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Invalid exchange type")); - } - - @Test - void createExchangeRequiresName() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 21, "method": "mq_create_exchange", - "params": { "name": " ", "type": "direct" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("name is required")); - } - - @Test - void deleteExchangeRejectsDefaultAndBuiltInsBeforeConnecting() { - String defaultResponse = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 22, "method": "mq_delete_exchange", "params": { "name": "" } } - """); - assertTrue(JsonParser.parseString(defaultResponse).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString().contains("default exchange cannot be deleted")); - - String builtInResponse = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 23, "method": "mq_delete_exchange", "params": { "name": "amq.direct" } } - """); - assertTrue(JsonParser.parseString(builtInResponse).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString().contains("built-in exchange 'amq.direct' cannot be deleted")); - } - - @Test - void listExchangesRequiresConnection() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 24, "method": "mq_list_exchanges", "params": {} } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Not connected")); - } - - @Test - void listBindingsRequiresConnection() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 25, "method": "mq_list_bindings", "params": { "queue": "q1" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Not connected")); - } - - @Test - void bindRequiresSourceAndDestination() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 26, "method": "mq_bind", - "params": { "source": "", "destination": "q1", "destinationType": "queue" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("source is required")); - } - - @Test - void bindRejectsUnknownDestinationType() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 27, "method": "mq_bind", - "params": { "source": "ex1", "destination": "q1", "destinationType": "stream" } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("destinationType must be 'queue' or 'exchange'")); - } - - // ------------------------------------------------------------------- - // Client connections & channels - // ------------------------------------------------------------------- - - @Test - void mapsClientConnectionInfoWithRatesAndConnectedAt() { - var connection = RabbitMqAgent.clientConnectionInfoFromJson(JsonParser.parseString(""" - { - "name": "10.0.0.1:52364 -> 10.0.0.2:5672", - "user": "spring", - "peer_host": "10.0.0.1", - "peer_port": 52364, - "state": "running", - "channels": 3, - "vhost": "/", - "recv_oct_details": { "rate": 12.5 }, - "send_oct_details": { "rate": 0.0 }, - "connected_at": 1751900000000 - } - """).getAsJsonObject()); - assertEquals("10.0.0.1:52364 -> 10.0.0.2:5672", connection.get("name")); - assertEquals("spring", connection.get("user")); - assertEquals("10.0.0.1", connection.get("peerHost")); - assertEquals(52364L, connection.get("peerPort")); - assertEquals("running", connection.get("state")); - assertEquals(3L, connection.get("channels")); - assertEquals(12.5, (Double) connection.get("recvRate"), 0.0001); - assertEquals(0.0, (Double) connection.get("sendRate"), 0.0001); - assertEquals(1751900000000L, connection.get("connectedAt")); - } - - @Test - void clientConnectionInfoOmitsMissingRatesAndConnectedAt() { - var connection = RabbitMqAgent.clientConnectionInfoFromJson(JsonParser.parseString(""" - { - "name": "c1", - "user": "guest", - "peer_host": "10.0.0.1", - "peer_port": 1, - "state": "blocked", - "channels": 0 - } - """).getAsJsonObject()); - assertEquals("c1", connection.get("name")); - assertFalse(connection.containsKey("recvRate")); - assertFalse(connection.containsKey("sendRate")); - assertFalse(connection.containsKey("connectedAt")); - } - - @Test - void mapsChannelInfoToCamelCase() { - var channel = RabbitMqAgent.channelInfoFromJson(JsonParser.parseString(""" - { - "name": "10.0.0.1:52364 -> 10.0.0.2:5672 (1)", - "connection_details": { "name": "10.0.0.1:52364 -> 10.0.0.2:5672" }, - "state": "running", - "prefetch_count": 20, - "messages_unacknowledged": 4, - "consumer_count": 2 - } - """).getAsJsonObject()); - assertEquals("10.0.0.1:52364 -> 10.0.0.2:5672 (1)", channel.get("name")); - assertEquals("10.0.0.1:52364 -> 10.0.0.2:5672", channel.get("connectionName")); - assertEquals("running", channel.get("state")); - assertEquals(20, channel.get("prefetch")); - assertEquals(4L, channel.get("messagesUnacked")); - assertEquals(2L, channel.get("consumerCount")); - } - - @Test - void channelInfoOmitsMissingOptionalFields() { - var channel = RabbitMqAgent.channelInfoFromJson(JsonParser.parseString(""" - { "name": "c (1)", "state": "running" } - """).getAsJsonObject()); - assertFalse(channel.containsKey("connectionName")); - assertFalse(channel.containsKey("prefetch")); - assertFalse(channel.containsKey("messagesUnacked")); - assertFalse(channel.containsKey("consumerCount")); - } - - @Test - void channelMatchesConnectionByDetailsNameOrNamePrefix() { - var channel = RabbitMqAgent.channelInfoFromJson(JsonParser.parseString(""" - { - "name": "10.0.0.1:52364 -> 10.0.0.2:5672 (1)", - "connection_details": { "name": "10.0.0.1:52364 -> 10.0.0.2:5672" } - } - """).getAsJsonObject()); - assertTrue(RabbitMqAgent.channelMatchesConnection(channel, "10.0.0.1:52364 -> 10.0.0.2:5672")); - // Prefix match works even without connection_details. - var noDetails = RabbitMqAgent.channelInfoFromJson(JsonParser.parseString(""" - { "name": "10.0.0.1:52364 -> 10.0.0.2:5672 (1)" } - """).getAsJsonObject()); - assertTrue(RabbitMqAgent.channelMatchesConnection(noDetails, "10.0.0.1:52364 -> 10.0.0.2:5672")); - assertFalse(RabbitMqAgent.channelMatchesConnection(channel, "10.0.0.9:11111 -> 10.0.0.2:5672")); - assertFalse(RabbitMqAgent.channelMatchesConnection(noDetails, "10.0.0.9:11111 -> 10.0.0.2:5672")); - } - - @Test - void urlEncodeNameEncodesSpacesAndArrows() { - assertEquals("10.0.0.1%3A52364%20-%3E%2010.0.0.2%3A5672", - RabbitMqAgent.urlEncodeName("10.0.0.1:52364 -> 10.0.0.2:5672")); - assertEquals("plain-name", RabbitMqAgent.urlEncodeName("plain-name")); - } - - @Test - void listConnectionsRequiresConnection() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 30, "method": "mq_list_connections", "params": {} } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Not connected")); - } - - @Test - void listChannelsRequiresConnection() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 31, "method": "mq_list_channels", "params": {} } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("Not connected")); - } - - @Test - void closeConnectionRequiresName() { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 32, "method": "mq_close_connection", "params": { "name": " " } } - """); - JsonObject error = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error"); - assertEquals(-1, error.get("code").getAsInt()); - assertTrue(error.get("message").getAsString().contains("name is required")); - } - - // ------------------------------------------------------------------- - // AMQP error mapping - // ------------------------------------------------------------------- - - @Test - void mapsResourceLockedToExclusiveQueueHint() { - String message = RabbitMqAgent.mapAmqpError(405, - "RESOURCE_LOCKED - cannot obtain exclusive access to locked queue " - + "'springCloudBus.anonymous.abc' in vhost '/'"); - assertTrue(message.contains("Queue 'springCloudBus.anonymous.abc' is exclusive")); - assertTrue(message.contains("owned by another connection")); - assertTrue(message.contains("Hint:")); - assertTrue(message.contains("management API")); - } - - @Test - void mapsResourceLockedWithoutQueueName() { - String message = RabbitMqAgent.mapAmqpError(405, "RESOURCE_LOCKED"); - assertTrue(message.startsWith("The queue is exclusive")); - assertTrue(message.contains("Hint:")); - } - - @Test - void mapsNotFoundToFriendlyQueueMessage() { - String message = RabbitMqAgent.mapAmqpError(404, "NOT_FOUND - no queue 'gone' in vhost '/'"); - assertEquals("Queue 'gone' was not found." - + " Hint: it may have been deleted, or it never existed on this virtual host.", message); - } - - @Test - void mapsNotFoundForExchange() { - String message = RabbitMqAgent.mapAmqpError(404, "NOT_FOUND - no exchange 'ex1' in vhost '/'"); - assertTrue(message.startsWith("Exchange 'ex1' was not found.")); - } - - @Test - void mapsPreconditionFailedToImmutableParametersHint() { - String message = RabbitMqAgent.mapAmqpError(406, - "PRECONDITION_FAILED - inequivalent arg 'durable' for queue 'q1' in vhost '/':" - + " received 'false' but current is 'true'"); - // The resource name is the queue, not the mismatched argument. - assertTrue(message.contains("Queue 'q1' already exists with different parameters.")); - assertFalse(message.contains("'durable' already exists")); - assertTrue(message.contains("Hint:")); - assertTrue(message.contains("immutable")); - assertTrue(message.contains("delete and re-declare")); - } - - @Test - void mapsPreconditionFailedForExchange() { - String message = RabbitMqAgent.mapAmqpError(406, - "PRECONDITION_FAILED - inequivalent arg 'type' for exchange 'ex1' in vhost '/':" - + " received 'fanout' but current is 'direct'"); - assertTrue(message.startsWith("Exchange 'ex1' already exists with different parameters.")); - assertTrue(message.contains("delete and re-declare the exchange")); - } - - @Test - void mapsPreconditionFailedWithoutResourceName() { - String message = RabbitMqAgent.mapAmqpError(406, "PRECONDITION_FAILED"); - assertTrue(message.startsWith("The queue already exists with different parameters.")); - assertTrue(message.contains("Hint:")); - } - - @Test - void mapsAccessRefusedToPermissionHint() { - String message = RabbitMqAgent.mapAmqpError(403, - "ACCESS_REFUSED - access to queue 'q1' in vhost '/' refused for user 'dbx'"); - assertTrue(message.contains("Access to 'q1' was refused.")); - assertTrue(message.contains("Hint:")); - assertTrue(message.contains("configure/write/read permissions")); - } - - @Test - void mapsAccessRefusedWithoutResourceName() { - String message = RabbitMqAgent.mapAmqpError(403, "ACCESS_REFUSED"); - assertTrue(message.startsWith("Access to the requested resource was refused.")); - assertTrue(message.contains("Hint:")); - } - - @Test - void leavesOtherReplyCodesUnmapped() { - assertNull(RabbitMqAgent.mapAmqpError(503, "COMMAND_INVALID")); - assertNull(RabbitMqAgent.mapAmqpError(501, "FRAME_ERROR")); - } - - @Test - void extractsDeclaredResourceNameFromPreconditionFailedText() { - assertEquals("q1", RabbitMqAgent.extractDeclaredResourceName( - "PRECONDITION_FAILED - inequivalent arg 'durable' for queue 'q1' in vhost '/'")); - assertEquals("ex1", RabbitMqAgent.extractDeclaredResourceName( - "PRECONDITION_FAILED - inequivalent arg 'type' for exchange 'ex1' in vhost '/'")); - assertNull(RabbitMqAgent.extractDeclaredResourceName("no resource here")); - assertNull(RabbitMqAgent.extractDeclaredResourceName(null)); - } - - @Test - void extractsFirstQuotedNameFromReplyText() { - assertEquals("q1", RabbitMqAgent.extractQuotedName("NOT_FOUND - no queue 'q1' in vhost '/'")); - assertNull(RabbitMqAgent.extractQuotedName("no quoted name here")); - assertNull(RabbitMqAgent.extractQuotedName(null)); - } - - @Test - void normalizeErrorMessageUsesAmqpFriendlyMapping() { - ShutdownSignalException shutdown = new ShutdownSignalException(true, false, - new AMQImpl.Channel.Close(405, - "RESOURCE_LOCKED - cannot obtain exclusive access to locked queue 'q1' in vhost '/'", - 0, 0), null); - String message = RabbitMqAgent.normalizeErrorMessage( - new java.io.IOException("channel is already closed", shutdown)); - assertTrue(message.contains("Queue 'q1' is exclusive")); - assertFalse(message.contains("RESOURCE_LOCKED")); - } - - @Test - void normalizeErrorMessageKeepsRawMessageForUnmappedCodes() { - ShutdownSignalException shutdown = new ShutdownSignalException(true, false, - new AMQImpl.Channel.Close(503, "COMMAND_INVALID - unknown method", 0, 0), null); - String message = RabbitMqAgent.normalizeErrorMessage(new java.io.IOException("boom", shutdown)); - assertTrue(message.contains("COMMAND_INVALID")); - } - - @Test - void normalizeErrorMessagePassesThroughPlainExceptions() { - String message = RabbitMqAgent.normalizeErrorMessage(new IllegalStateException("Not connected")); - assertEquals("Not connected", message); - } - - // ------------------------------------------------------------------- - // Users & permissions - // ------------------------------------------------------------------- - - @Test - void mapsUserInfoWithTagsArray() { - var user = RabbitMqAgent.userInfoFromJson(JsonParser.parseString(""" - { "name": "jjsd", "tags": "administrator,management" } - """).getAsJsonObject()); - assertEquals("jjsd", user.get("name")); - assertEquals(List.of("administrator", "management"), user.get("tags")); - } - - @Test - void parseUserTagsTrimsAndDropsBlanks() { - assertEquals(List.of("administrator", "monitoring"), - RabbitMqAgent.parseUserTags("administrator, monitoring,,")); - assertTrue(RabbitMqAgent.parseUserTags("").isEmpty()); - assertTrue(RabbitMqAgent.parseUserTags(" , ").isEmpty()); - } - - @Test - void userTagsParamAcceptsArrayOrString() { - assertEquals("management,policymaker", RabbitMqAgent.userTagsParam(JsonParser.parseString(""" - { "tags": ["management", " policymaker "] } - """).getAsJsonObject())); - assertEquals("administrator", RabbitMqAgent.userTagsParam(JsonParser.parseString(""" - { "tags": "administrator" } - """).getAsJsonObject())); - assertEquals("", RabbitMqAgent.userTagsParam(JsonParser.parseString(""" - { "name": "dbx-test-user" } - """).getAsJsonObject())); - } - - @Test - void mapsPermissionInfo() { - var permission = RabbitMqAgent.permissionInfoFromJson(JsonParser.parseString(""" - { "user": "jjsd", "vhost": "/", "configure": ".*", "write": ".*", "read": ".*" } - """).getAsJsonObject()); - assertEquals("jjsd", permission.get("user")); - assertEquals("/", permission.get("vhost")); - assertEquals(".*", permission.get("configure")); - assertEquals(".*", permission.get("write")); - assertEquals(".*", permission.get("read")); - } - - @Test - void permissionPatternDefaultsToMatchAll() { - JsonObject params = JsonParser.parseString(""" - { "write": "^dbx-", "read": "" } - """).getAsJsonObject(); - assertEquals(".*", RabbitMqAgent.permissionPattern(params, "configure")); - assertEquals("^dbx-", RabbitMqAgent.permissionPattern(params, "write")); - assertEquals(".*", RabbitMqAgent.permissionPattern(params, "read")); - } - - @Test - void permissionVhostRejectsBlankAndAllVhostsSentinel() { - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.permissionVhost( - JsonParser.parseString("{}").getAsJsonObject())); - assertThrows(IllegalArgumentException.class, () -> RabbitMqAgent.permissionVhost( - JsonParser.parseString(""" - { "virtual_host": "*" } - """).getAsJsonObject())); - assertEquals("/", RabbitMqAgent.permissionVhost(JsonParser.parseString(""" - { "virtual_host": "/" } - """).getAsJsonObject())); - } - - @Test - void userGuardRejectsConnectedUser() { - assertThrows(IllegalArgumentException.class, - () -> RabbitMqAgent.assertNotConnectedUser("delete", "jjsd", "jjsd")); - assertThrows(IllegalArgumentException.class, - () -> RabbitMqAgent.assertNotConnectedUser("create or modify", "jjsd", "jjsd")); - RabbitMqAgent.assertNotConnectedUser("delete", "dbx-test-user", "jjsd"); - } - - @Test - void createUserRequiresNameAndPassword() { - String noName = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 50, "method": "mq_create_user", "params": { "password": "x" } } - """); - assertTrue(JsonParser.parseString(noName).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString().contains("user name is required")); - - String noPassword = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 51, "method": "mq_create_user", "params": { "name": "dbx-test-user" } } - """); - assertTrue(JsonParser.parseString(noPassword).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString().contains("password is required")); - } - - @Test - void grantAndRevokeRequireVhostBeforeConnecting() { - String grant = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 52, "method": "mq_grant_permission", "params": { "user": "dbx-test-user" } } - """); - String grantMessage = JsonParser.parseString(grant).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString(); - assertTrue(grantMessage.contains("virtual_host is required"), grantMessage); - - String revoke = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 53, "method": "mq_revoke_permission", - "params": { "user": "dbx-test-user", "virtual_host": "*" } } - """); - String revokeMessage = JsonParser.parseString(revoke).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString(); - assertTrue(revokeMessage.contains("all_vhosts is only supported for list operations"), revokeMessage); - } - - @Test - void grantAndRevokeRejectAllVhostsFlag() { - for (String method : List.of("mq_grant_permission", "mq_revoke_permission")) { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 54, "method": "%s", - "params": { "all_vhosts": true, "user": "dbx-test-user", "virtual_host": "/" } } - """.formatted(method))).getAsJsonObject(); - assertEquals("all_vhosts is only supported for list operations", - response.getAsJsonObject("error").get("message").getAsString(), method); - } - } - - @Test - void userAndPermissionOperationsRequireConnection() { - List requests = List.of(""" - { "jsonrpc": "2.0", "id": 55, "method": "mq_list_users", "params": {} } - """, """ - { "jsonrpc": "2.0", "id": 56, "method": "mq_list_permissions", "params": {} } - """, """ - { "jsonrpc": "2.0", "id": 57, "method": "mq_delete_user", "params": { "name": "dbx-test-user" } } - """); - for (String request : requests) { - String message = JsonParser.parseString(RabbitMqAgent.handleRequest(request)) - .getAsJsonObject().getAsJsonObject("error").get("message").getAsString(); - assertTrue(message.contains("Not connected"), message); - } - } - - @Test - void deleteUserRejectsConnectedUserBeforeHttpCall() { - // The guard must win over the management API call: no broker is needed. - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 58, "method": "mq_delete_user", - "params": { "name": "jjsd", - "connection": { "addresses": "127.0.0.1:1", "username": "jjsd" } } } - """); - String message = JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString(); - assertTrue(message.contains("Cannot delete user 'jjsd' while connected as that user"), message); - } - - // ------------------------------------------------------------------- - // Policies - // ------------------------------------------------------------------- - - @Test - void policyInfoMapsApplyToAndDefinition() { - Map policy = RabbitMqAgent.policyInfoFromJson(JsonParser.parseString(""" - { "name": "dbx-pol", "vhost": "/", "pattern": "^dbx-", - "apply-to": "exchanges", "priority": 5, - "definition": { "max-length": 100, "alternate-exchange": "dbx-ae", "skip": null } } - """).getAsJsonObject()); - assertEquals("dbx-pol", policy.get("name")); - assertEquals("/", policy.get("vhost")); - assertEquals("^dbx-", policy.get("pattern")); - assertEquals("exchanges", policy.get("applyTo")); - assertEquals(5L, policy.get("priority")); - @SuppressWarnings("unchecked") - Map definition = (Map) policy.get("definition"); - assertEquals(100L, definition.get("max-length")); - assertEquals("dbx-ae", definition.get("alternate-exchange")); - assertFalse(definition.containsKey("skip")); - } - - @Test - void listPoliciesMapsEntriesViaManagementApi() throws Exception { - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/api/policies", exchange -> { - byte[] body = """ - [ { "name": "dbx-pol", "vhost": "/", "pattern": "^dbx-", - "apply-to": "queues", "priority": 0, - "definition": { "max-length": 100 } } ] - """.getBytes(StandardCharsets.UTF_8); - exchange.getResponseHeaders().add("Content-Type", "application/json"); - exchange.sendResponseHeaders(200, body.length); - exchange.getResponseBody().write(body); - exchange.close(); - }); - server.start(); - try { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 60, "method": "mq_list_policies", - "params": { "all_vhosts": true, - "connection": { "addresses": "127.0.0.1", - "properties": { "management_port": %d } } } } - """.formatted(server.getAddress().getPort()))).getAsJsonObject(); - JsonObject policy = response.getAsJsonObject("result").getAsJsonArray("policies") - .get(0).getAsJsonObject(); - assertEquals("dbx-pol", policy.get("name").getAsString()); - assertEquals("/", policy.get("vhost").getAsString()); - assertEquals("queues", policy.get("applyTo").getAsString()); - assertEquals(100, policy.getAsJsonObject("definition").get("max-length").getAsInt()); - } finally { - server.stop(0); - } - } - - @Test - void setPolicyAppliesDefaultsAndMapsApplyTo() throws Exception { - String[] capturedRequest = new String[2]; // [0] "METHOD path", [1] request body - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/", exchange -> { - capturedRequest[0] = exchange.getRequestMethod() + " " + exchange.getRequestURI().getRawPath(); - capturedRequest[1] = new String(exchange.getRequestBody().readAllBytes(), StandardCharsets.UTF_8); - exchange.sendResponseHeaders(204, -1); - exchange.close(); - }); - server.start(); - try { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 61, "method": "mq_set_policy", - "params": { "virtual_host": "/", "name": "dbx-pol", "pattern": "^dbx-", - "definition": { "max-length": 100 }, - "connection": { "addresses": "127.0.0.1", - "properties": { "management_port": %d } } } } - """.formatted(server.getAddress().getPort())); - assertTrue(JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("result") - .get("ok").getAsBoolean()); - assertEquals("PUT /api/policies/%2F/dbx-pol", capturedRequest[0]); - JsonObject body = JsonParser.parseString(capturedRequest[1]).getAsJsonObject(); - // applyTo defaults to queues and priority to 0. - assertEquals("queues", body.get("apply-to").getAsString()); - assertEquals(0, body.get("priority").getAsInt()); - assertEquals("^dbx-", body.get("pattern").getAsString()); - assertEquals(100, body.getAsJsonObject("definition").get("max-length").getAsInt()); - } finally { - server.stop(0); - } - } - - @Test - void deletePolicyCallsManagementApiDelete() throws Exception { - String[] capturedRequest = new String[1]; - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/", exchange -> { - capturedRequest[0] = exchange.getRequestMethod() + " " + exchange.getRequestURI().getRawPath(); - exchange.sendResponseHeaders(204, -1); - exchange.close(); - }); - server.start(); - try { - String response = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 62, "method": "mq_delete_policy", - "params": { "virtual_host": "/", "name": "dbx-pol", - "connection": { "addresses": "127.0.0.1", - "properties": { "management_port": %d } } } } - """.formatted(server.getAddress().getPort())); - assertTrue(JsonParser.parseString(response).getAsJsonObject().getAsJsonObject("result") - .get("ok").getAsBoolean()); - assertEquals("DELETE /api/policies/%2F/dbx-pol", capturedRequest[0]); - } finally { - server.stop(0); - } - } - - @Test - void setAndDeletePolicyRejectAllVhostsSentinel() { - for (String method : List.of("mq_set_policy", "mq_delete_policy")) { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 63, "method": "%s", - "params": { "virtual_host": "*", "name": "dbx-pol", "pattern": "^dbx-", - "definition": {} } } - """.formatted(method))).getAsJsonObject(); - assertEquals("all_vhosts is only supported for list operations", - response.getAsJsonObject("error").get("message").getAsString(), method); - } - } - - @Test - void setAndDeletePolicyRejectAllVhostsFlag() { - for (String method : List.of("mq_set_policy", "mq_delete_policy")) { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 64, "method": "%s", - "params": { "all_vhosts": true, "virtual_host": "/", "name": "dbx-pol" } } - """.formatted(method))).getAsJsonObject(); - assertEquals("all_vhosts is only supported for list operations", - response.getAsJsonObject("error").get("message").getAsString(), method); - } - } - - @Test - void setPolicyRequiresNamePatternAndDefinition() { - String noName = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 65, "method": "mq_set_policy", - "params": { "virtual_host": "/" } } - """); - assertTrue(JsonParser.parseString(noName).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString().contains("name is required")); - - String noPattern = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 66, "method": "mq_set_policy", - "params": { "virtual_host": "/", "name": "dbx-pol", "definition": {} } } - """); - assertTrue(JsonParser.parseString(noPattern).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString().contains("pattern is required")); - - String noDefinition = RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 67, "method": "mq_set_policy", - "params": { "virtual_host": "/", "name": "dbx-pol", "pattern": "^dbx-" } } - """); - assertTrue(JsonParser.parseString(noDefinition).getAsJsonObject().getAsJsonObject("error") - .get("message").getAsString().contains("definition is required")); - } - - // ------------------------------------------------------------------- - // Overview & nodes - // ------------------------------------------------------------------- - - @Test - void overviewInfoMapsTotalsAndRates() { - Map overview = RabbitMqAgent.overviewInfoFromJson(JsonParser.parseString(""" - { "queue_totals": { "messages_ready": 12, "messages_unacknowledged": 3 }, - "message_stats": { "publish": 100, "publish_details": { "rate": 1.5 }, - "deliver_get": 90, "deliver_get_details": { "rate": 2.5 }, - "ack": 80, "ack_details": { "rate": 0.5 } }, - "object_totals": { "connections": 4, "channels": 6, "exchanges": 8, - "queues": 10, "consumers": 2 } } - """).getAsJsonObject()); - assertEquals(12L, overview.get("messagesReady")); - assertEquals(3L, overview.get("messagesUnacked")); - assertEquals(1.5, (Double) overview.get("publishRate"), 0.0001); - assertEquals(2.5, (Double) overview.get("deliverRate"), 0.0001); - assertEquals(0.5, (Double) overview.get("ackRate"), 0.0001); - assertEquals(10L, overview.get("totalQueues")); - assertEquals(8L, overview.get("totalExchanges")); - assertEquals(4L, overview.get("totalConnections")); - assertEquals(6L, overview.get("totalChannels")); - assertEquals(2L, overview.get("totalConsumers")); - } - - @Test - void overviewInfoOmitsMissingStats() { - Map overview = RabbitMqAgent.overviewInfoFromJson(JsonParser.parseString(""" - { "queue_totals": { "messages_ready": 1 } } - """).getAsJsonObject()); - assertEquals(1L, overview.get("messagesReady")); - assertFalse(overview.containsKey("messagesUnacked")); - assertFalse(overview.containsKey("publishRate")); - assertFalse(overview.containsKey("totalQueues")); - } - - @Test - void nodeInfoMapsSnakeCaseToCamelCase() { - Map node = RabbitMqAgent.nodeInfoFromJson(JsonParser.parseString(""" - { "name": "rabbit@node1", "running": true, "mem_used": 1000, "mem_limit": 2000, - "disk_free": 3000, "fd_used": 10, "fd_total": 100, "sockets_used": 5, - "sockets_total": 50, "uptime": 123456 } - """).getAsJsonObject()); - assertEquals("rabbit@node1", node.get("name")); - assertEquals(true, node.get("running")); - assertEquals(1000L, node.get("memUsed")); - assertEquals(2000L, node.get("memLimit")); - assertEquals(3000L, node.get("diskFree")); - assertEquals(10L, node.get("fdUsed")); - assertEquals(100L, node.get("fdTotal")); - assertEquals(5L, node.get("socketsUsed")); - assertEquals(50L, node.get("socketsTotal")); - assertEquals(123456L, node.get("uptimeMs")); - } - - @Test - void listNodesMapsEntriesViaManagementApi() throws Exception { - HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); - server.createContext("/api/nodes", exchange -> { - byte[] body = """ - [ { "name": "rabbit@node1", "running": true, "mem_used": 1000, - "uptime": 123456 } ] - """.getBytes(StandardCharsets.UTF_8); - exchange.getResponseHeaders().add("Content-Type", "application/json"); - exchange.sendResponseHeaders(200, body.length); - exchange.getResponseBody().write(body); - exchange.close(); - }); - server.start(); - try { - JsonObject response = JsonParser.parseString(RabbitMqAgent.handleRequest(""" - { "jsonrpc": "2.0", "id": 68, "method": "mq_list_nodes", - "params": { "connection": { "addresses": "127.0.0.1", - "properties": { "management_port": %d } } } } - """.formatted(server.getAddress().getPort()))).getAsJsonObject(); - JsonObject node = response.getAsJsonObject("result").getAsJsonArray("nodes") - .get(0).getAsJsonObject(); - assertEquals("rabbit@node1", node.get("name").getAsString()); - assertTrue(node.get("running").getAsBoolean()); - assertEquals(1000, node.get("memUsed").getAsLong()); - assertEquals(123456, node.get("uptimeMs").getAsLong()); - } finally { - server.stop(0); - } - } - - // ------------------------------------------------------------------- - // Channel self-healing decision - // ------------------------------------------------------------------- - - @Test - void nullChannelNeedsRecreation() { - assertTrue(RabbitMqAgent.needsNewChannel(null)); - } - - @Test - void closedChannelNeedsRecreation() { - assertTrue(RabbitMqAgent.needsNewChannel(stubChannel(false))); - } - - @Test - void openChannelIsReused() { - assertFalse(RabbitMqAgent.needsNewChannel(stubChannel(true))); - } - - private static Channel stubChannel(boolean open) { - return (Channel) Proxy.newProxyInstance( - RabbitMqAgentTest.class.getClassLoader(), - new Class[] { Channel.class }, - (proxy, method, args) -> { - if ("isOpen".equals(method.getName())) { - return open; - } - throw new UnsupportedOperationException(method.getName()); - }); - } -} diff --git a/agents/scripts/driver_release_packages_test.py b/agents/scripts/driver_release_packages_test.py index 3ecd33398..148eee9ef 100644 --- a/agents/scripts/driver_release_packages_test.py +++ b/agents/scripts/driver_release_packages_test.py @@ -20,6 +20,8 @@ class DriverReleasePackagesTest(unittest.TestCase): native_source.write_bytes(b"MZtest-agent") duckdb_source = release_dir / "dbx-agent-duckdb-macos-aarch64" duckdb_source.write_bytes(b"\xcf\xfa\xed\xfetest-duckdb-agent") + rabbitmq_source = release_dir / "dbx-agent-rabbitmq-linux-x64" + rabbitmq_source.write_bytes(b"\x7fELFtest-rabbitmq-agent") java_source = release_dir / "dbx-agent-h2.jar" java_source.write_bytes(b"test-jar") versions = { @@ -28,13 +30,15 @@ class DriverReleasePackagesTest(unittest.TestCase): "xugu": "0.1.20", "kingbase": "0.1.34", "duckdb": "0.1.0", + "rabbitmq": "0.1.0", } renamed = version_agent_artifacts(release_dir, versions) versioned_java = release_dir / "dbx-agent-h2-0.2.5.jar" versioned_native = release_dir / "dbx-agent-kingbase-0.1.34-windows-x64.exe" versioned_duckdb = release_dir / "dbx-agent-duckdb-0.1.0-macos-aarch64" - self.assertEqual(renamed, [versioned_java, versioned_native, versioned_duckdb]) + versioned_rabbitmq = release_dir / "dbx-agent-rabbitmq-0.1.0-linux-x64" + self.assertEqual(renamed, [versioned_java, versioned_native, versioned_duckdb, versioned_rabbitmq]) registry = { "jres": {"21": {"version": "21", "platforms": {}}}, @@ -72,6 +76,19 @@ class DriverReleasePackagesTest(unittest.TestCase): } }, }, + "rabbitmq": { + "version": "0.1.0", + "label": "RabbitMQ", + "min_app_version": "0.6.0", + "jre": "21", + "jar": {"url": "https://example.com/legacy-placeholder.jar", "size": 0}, + "native": { + "linux-x64": { + "url": f"https://example.com/{versioned_rabbitmq.name}", + "size": versioned_rabbitmq.stat().st_size, + } + }, + }, }, } (release_dir / "agent-registry.json").write_text(json.dumps(registry), encoding="utf-8") @@ -84,12 +101,14 @@ class DriverReleasePackagesTest(unittest.TestCase): release_dir / "dbx-agent-h2-0.2.5.tar.zst", release_dir / "dbx-agent-kingbase-0.1.34-windows-x64.tar.zst", release_dir / "dbx-agent-duckdb-0.1.0-macos-aarch64.tar.zst", + release_dir / "dbx-agent-rabbitmq-0.1.0-linux-x64.tar.zst", ], ) package_cases = [ (outputs[0], "h2", versioned_java, "jar", None), (outputs[1], "kingbase", versioned_native, "native", "windows-x64"), (outputs[2], "duckdb", versioned_duckdb, "native", "macos-aarch64"), + (outputs[3], "rabbitmq", versioned_rabbitmq, "native", "linux-x64"), ] for output, driver_name, source, artifact_type, platform in package_cases: tar_bytes = subprocess.run( @@ -116,6 +135,7 @@ class DriverReleasePackagesTest(unittest.TestCase): (final_registry["drivers"]["h2"]["jar"], outputs[0]), (final_registry["drivers"]["kingbase"]["native"]["windows-x64"], outputs[1]), (final_registry["drivers"]["duckdb"]["native"]["macos-aarch64"], outputs[2]), + (final_registry["drivers"]["rabbitmq"]["native"]["linux-x64"], outputs[3]), ] for artifact, output in release_artifacts: self.assertEqual(artifact["url"], f"https://example.com/{output.name}") @@ -124,7 +144,7 @@ class DriverReleasePackagesTest(unittest.TestCase): self.assertEqual(len(artifact["sha256"]), 64) removed = remove_raw_driver_artifacts(release_dir) - self.assertEqual(removed, [versioned_duckdb, versioned_java, versioned_native]) + self.assertEqual(removed, [versioned_duckdb, versioned_java, versioned_native, versioned_rabbitmq]) self.assertTrue(all(output.is_file() for output in outputs)) def test_full_offline_bundle_includes_supported_windows_artifacts(self) -> None: diff --git a/agents/scripts/validate_agents.py b/agents/scripts/validate_agents.py index fc030bdc4..02cc61d52 100644 --- a/agents/scripts/validate_agents.py +++ b/agents/scripts/validate_agents.py @@ -17,6 +17,7 @@ NATIVE_ONLY_AGENT_MODULES = { "oracle": "drivers/oracle-go", "kingbase": "drivers/kingbase-go", "xugu": "drivers/xugu", + "rabbitmq": "drivers/rabbitmq", } AUTO_VERSIONED_NATIVE_MODULES = {"duckdb"} JDBC_ARCHITECTURE_ALLOWLIST = { diff --git a/agents/scripts/validate_agents_test.py b/agents/scripts/validate_agents_test.py index 8c438d06c..e1be4bfdb 100644 --- a/agents/scripts/validate_agents_test.py +++ b/agents/scripts/validate_agents_test.py @@ -136,10 +136,10 @@ class ValidateAgentsTest(unittest.TestCase): "include(*(infrastructureModules + driverModules))\n", encoding="utf-8", ) - for driver in ("oracle-go", "kingbase-go", "xugu", "duckdb"): + for driver in ("oracle-go", "kingbase-go", "xugu", "duckdb", "rabbitmq"): (root / "drivers" / driver).mkdir(parents=True) (root / "versions.json").write_text( - json.dumps({"h2": "0.1.0", "oracle": "0.1.0", "kingbase": "0.1.0", "xugu": "0.1.0"}), + json.dumps({"h2": "0.1.0", "oracle": "0.1.0", "kingbase": "0.1.0", "xugu": "0.1.0", "rabbitmq": "0.1.0"}), encoding="utf-8", ) diff --git a/agents/scripts/version_agent_artifacts.py b/agents/scripts/version_agent_artifacts.py index f02e2f88f..ebbce01cd 100644 --- a/agents/scripts/version_agent_artifacts.py +++ b/agents/scripts/version_agent_artifacts.py @@ -4,7 +4,7 @@ import json from pathlib import Path -NATIVE_DRIVERS = ("oracle", "xugu", "kingbase", "duckdb") +NATIVE_DRIVERS = ("oracle", "xugu", "kingbase", "duckdb", "rabbitmq") PLATFORMS = ( "macos-aarch64", "macos-x64", diff --git a/agents/settings.gradle b/agents/settings.gradle index b06bec2f0..12e530ea1 100644 --- a/agents/settings.gradle +++ b/agents/settings.gradle @@ -6,7 +6,7 @@ def driverModules = [ 'teradata', 'vertica', 'firebird', 'exasol', 'oceanbase-oracle', 'gbase8a', 'gbase8s', 'bigquery', 'kylin', 'sundb', 'h2', 'h2-legacy', 'snowflake', 'trino', 'hive', 'spark', 'db2', 'informix', 'neo4j', 'cassandra', 'mongodb', 'highgo', 'uxdb', 'tdengine', 'yashandb', 'oscar', - 'iris', 'iotdb', 'etcd', 'zookeeper', 'kafka', 'rocketmq', 'rabbitmq', 'sqlserver-legacy' + 'iris', 'iotdb', 'etcd', 'zookeeper', 'kafka', 'rocketmq', 'sqlserver-legacy' ] include(*(infrastructureModules + driverModules)) diff --git a/crates/dbx-core/src/mq/README.md b/crates/dbx-core/src/mq/README.md index dc6fe1f7c..09ff7333a 100644 --- a/crates/dbx-core/src/mq/README.md +++ b/crates/dbx-core/src/mq/README.md @@ -14,7 +14,7 @@ mq/ ├── service.rs - 服务层函数 └── adapters/ ├── pulsar.rs - Pulsar 实现 - ├── rabbitmq.rs - RabbitMQ 实现 (Java agent) + ├── rabbitmq.rs - RabbitMQ 实现 (Go native agent) └── pulsar_version.rs - 版本探测 ``` diff --git a/crates/dbx-core/src/mq/adapters/rabbitmq.rs b/crates/dbx-core/src/mq/adapters/rabbitmq.rs index f58730fe5..fb2d2bbbf 100644 --- a/crates/dbx-core/src/mq/adapters/rabbitmq.rs +++ b/crates/dbx-core/src/mq/adapters/rabbitmq.rs @@ -1,9 +1,9 @@ -//! RabbitMQ admin adapter. Communicates with a Java agent process -//! (`RabbitMqAgent.java`) via JSON-RPC over stdin/stdout. The Java agent uses -//! the `amqp-client` library for admin and message operations. +//! RabbitMQ admin adapter. Communicates with the native Go agent via JSON-RPC +//! over stdin/stdout. The agent uses `amqp091-go` for AMQP operations and the +//! RabbitMQ management HTTP API for administrative operations. //! //! This adapter follows the same pattern as the Kafka agent: -//! 1. Spawn a Java agent process via `AgentDriverClient` +//! 1. Spawn the native agent process via `AgentDriverClient` //! 2. Perform JSON-RPC handshake + connect //! 3. Delegate all `MessageQueueAdmin` trait methods to JSON-RPC calls @@ -59,7 +59,7 @@ pub struct RabbitMqAdmin { } impl RabbitMqAdmin { - /// Spawn the RabbitMQ Java agent, perform handshake, and connect. + /// Spawn the RabbitMQ native agent, perform handshake, and connect. pub async fn new(cfg: MqAdminConfig, launch: AgentLaunchSpec) -> Result { let mut client = AgentDriverClient::spawn(launch).await?; diff --git a/docs/mq-quick-start.md b/docs/mq-quick-start.md index d6b81e1d3..eaced23db 100644 --- a/docs/mq-quick-start.md +++ b/docs/mq-quick-start.md @@ -444,8 +444,9 @@ Agent 构建与安装: ```bash cd agents -./gradlew :rabbitmq:shadowJar -# 将 shadow JAR 安装到 DBX 数据目录 agents/drivers/rabbitmq/agent.jar +cd drivers/rabbitmq +go build -o agent . +# 将原生 agent 安装到 DBX 数据目录 agents/drivers/rabbitmq/agent ``` Docker 快速启动(AMQP 5672 + Management 15672,仅用于本地验证): @@ -506,4 +507,3 @@ const response = await mqRawRequest(connectionId, { }) console.log(response.body) ``` -