diff --git a/.github/scripts/bump-agent-versions.mjs b/.github/scripts/bump-agent-versions.mjs index 7fb5fef8f9..76d0e81e56 100644 --- a/.github/scripts/bump-agent-versions.mjs +++ b/.github/scripts/bump-agent-versions.mjs @@ -60,9 +60,10 @@ const nativeDriverDirectories = { rabbitmq: "rabbitmq", rocketmq: "rocketmq", zookeeper: "zookeeper", + nats: "nats", tdengine: "tdengine", }; -const nativeDriverModules = new Set(["cassandra", "duckdb", "hive", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "rabbitmq", "rocketmq", "zookeeper", "tdengine"]); +const nativeDriverModules = new Set(["cassandra", "duckdb", "hive", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "rabbitmq", "rocketmq", "zookeeper", "nats", "tdengine"]); const nativeDriverSharedPaths = { hive: [ "agents/go-common/go-gssapi", diff --git a/.github/scripts/bump-agent-versions.test.mjs b/.github/scripts/bump-agent-versions.test.mjs index ea6e6e06cc..dcc6097907 100644 --- a/.github/scripts/bump-agent-versions.test.mjs +++ b/.github/scripts/bump-agent-versions.test.mjs @@ -90,6 +90,19 @@ test("bumps ZooKeeper from native and shared SASL source directories", () => { } }); +test("bumps the native NATS agent from its Go directory", () => { + const result = evaluateAgentVersionBump({ + versions: { nats: "0.1.0" }, + changedFiles: ["agents/drivers/nats/main.go"], + moduleExists: (path) => path === "agents/drivers/nats", + readModuleFile: () => "", + }); + + assert.equal(result.versions.nats, "0.1.1"); + assert.deepEqual(result.javaModules, []); + assert.deepEqual(result.nativeModules, ["nats"]); +}); + test("bumps the native Vastbase agent from its independent Go directory", () => { const result = evaluateAgentVersionBump({ versions: { vastbase: "0.1.37" }, diff --git a/.github/scripts/reuse-agent-release-assets.mjs b/.github/scripts/reuse-agent-release-assets.mjs index 61f9159e24..18b73005fc 100644 --- a/.github/scripts/reuse-agent-release-assets.mjs +++ b/.github/scripts/reuse-agent-release-assets.mjs @@ -15,7 +15,7 @@ import { basename, join } from "node:path"; import { tmpdir } from "node:os"; const REGISTRY_ASSET = "agent-registry.json"; -const NATIVE_MODULES = new Set(["duckdb", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "rabbitmq", "rocketmq", "zookeeper", "tdengine"]); +const NATIVE_MODULES = new Set(["duckdb", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "rabbitmq", "rocketmq", "zookeeper", "nats", "tdengine"]); const PLATFORMS = [ "macos-aarch64", "macos-x64", diff --git a/.github/scripts/reuse-agent-release-assets.test.mjs b/.github/scripts/reuse-agent-release-assets.test.mjs index 174a2865bb..49e93525e2 100644 --- a/.github/scripts/reuse-agent-release-assets.test.mjs +++ b/.github/scripts/reuse-agent-release-assets.test.mjs @@ -132,6 +132,24 @@ test("requires all RocketMQ native platforms when reusing a release", () => { ); }); +test("requires all NATS native platforms when reusing a release", () => { + const native = Object.fromEntries( + platforms.slice(0, -1).map((platform, index) => [platform, artifact(`dbx-agent-nats-0.1.40-${platform}.tar.zst`, String(index + 1))]), + ); + const registry = { drivers: { nats: { version: "0.1.40", native } }, jres: {} }; + + assert.throws( + () => collectReusableAssetPlan({ + registry, + release: releaseFor(Object.values(native)), + versions: { nats: "0.1.40" }, + modules: ["nats"], + reuseJre: false, + }), + /missing=windows-x64/, + ); +}); + test("ignores zero-size legacy JAR placeholders for native-only modules", () => { const native = Object.fromEntries( platforms.map((platform, index) => [platform, artifact(`dbx-agent-duckdb-0.1.2-${platform}.tar.zst`, String(index + 1))]), diff --git a/.github/workflows/agents-release.yml b/.github/workflows/agents-release.yml index 16aa1aaefa..9f5cea41a4 100644 --- a/.github/workflows/agents-release.yml +++ b/.github/workflows/agents-release.yml @@ -369,6 +369,46 @@ jobs: name: zookeeper-native path: "release-native/dbx-agent-zookeeper-*" + build-nats-native: + needs: [bump-versions] + if: ${{ contains(fromJSON(needs.bump-versions.outputs.native_modules), 'nats') }} + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.22.x" + - name: Test NATS native agent + working-directory: agents/drivers/nats + run: go test ./... + - name: Cross-compile NATS native agent + shell: bash + run: | + mkdir -p release-native + cd agents/drivers/nats + 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-nats-${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: nats-native + path: "release-native/dbx-agent-nats-*" + build-cassandra-native: needs: [bump-versions] if: ${{ contains(fromJSON(needs.bump-versions.outputs.native_modules), 'cassandra') }} @@ -948,7 +988,7 @@ jobs: retention-days: 1 release: - needs: [bump-versions, commit-versions, build-agents, build-oracle-native, build-xugu-native, build-rabbitmq-native, build-rocketmq-native, build-zookeeper-native, build-cassandra-native, build-hive-native, build-kingbase-native, build-vastbase-native, build-neo4j-native, build-iotdb-native, build-duckdb-native, build-tdengine-native, build-jre, reuse-previous-assets] + needs: [bump-versions, commit-versions, build-agents, build-oracle-native, build-xugu-native, build-rabbitmq-native, build-rocketmq-native, build-zookeeper-native, build-nats-native, build-cassandra-native, build-hive-native, build-kingbase-native, build-vastbase-native, build-neo4j-native, build-iotdb-native, build-duckdb-native, build-tdengine-native, build-jre, reuse-previous-assets] if: ${{ always() && !contains(needs.*.result, 'failure') && !contains(needs.*.result, 'cancelled') }} runs-on: ubuntu-latest steps: @@ -1077,6 +1117,7 @@ jobs: rabbitmq) echo "RabbitMQ" ;; rocketmq) echo "Apache RocketMQ" ;; zookeeper) echo "Apache ZooKeeper" ;; + nats) echo "NATS" ;; cassandra) echo "Apache Cassandra" ;; hive) echo "Apache Hive" ;; neo4j) echo "Neo4j" ;; @@ -1145,7 +1186,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 vastbase neo4j iotdb duckdb rabbitmq rocketmq zookeeper cassandra hive tdengine; do + for name in oracle xugu kingbase vastbase neo4j iotdb duckdb rabbitmq rocketmq zookeeper nats cassandra hive tdengine; do version=$(get_module_version "$name") [ -f "release/dbx-agent-${name}-${version}.jar" ] && continue native_json=$(generate_native_platforms "$name" "$version") diff --git a/agents/README.md b/agents/README.md index 82b5b7b7e2..001eacf414 100644 --- a/agents/README.md +++ b/agents/README.md @@ -46,11 +46,12 @@ Each agent runs as a standalone process and communicates with DBX via stdin/stdo | zookeeper | Apache ZooKeeper | go-zookeeper native agent | | rabbitmq | RabbitMQ | amqp091-go native agent | | rocketmq | Apache RocketMQ | rocketmq-admin-go native agent | +| nats | NATS | nats.go native agent | ## Multi-JRE Support -Most Java agents target JRE 21. Native agents, such as `cassandra`, `duckdb`, `hive`, `iotdb`, `oracle`, `kingbase`, `tdengine`, `xugu`, `rabbitmq`, `rocketmq`, and `zookeeper`, 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 `cassandra`, `duckdb`, `hive`, `iotdb`, `oracle`, `kingbase`, `tdengine`, `xugu`, `rabbitmq`, `rocketmq`, `zookeeper`, and `nats`, do not require a JRE. DBX downloads and manages the JRE 21 installation automatically for Java agents. ## JDBC Connection Pooling @@ -76,7 +77,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/cassandra-go` (Apache cassandra-gocql-driver), `drivers/duckdb`, `drivers/hive-go` (native HS2), `drivers/iotdb` (Apache IoTDB Go Client), `drivers/oracle-go` (go-ora), `drivers/kingbase-go` (gokb), `drivers/vastbase-go` (openGauss connector), `drivers/tdengine` (taos-connector-rust), `drivers/xugu`, `drivers/rabbitmq` (amqp091-go), `drivers/rocketmq` (rocketmq-admin-go), and `drivers/zookeeper` (go-zookeeper) as reference implementations. No JRE download or management is needed. +- **Native (C++/Go/Rust)** — preferred when a usable native driver exists. See `drivers/cassandra-go` (Apache cassandra-gocql-driver), `drivers/duckdb`, `drivers/hive-go` (native HS2), `drivers/iotdb` (Apache IoTDB Go Client), `drivers/oracle-go` (go-ora), `drivers/kingbase-go` (gokb), `drivers/vastbase-go` (openGauss connector), `drivers/tdengine` (taos-connector-rust), `drivers/xugu`, `drivers/rabbitmq` (amqp091-go), `drivers/rocketmq` (rocketmq-admin-go), `drivers/zookeeper` (go-zookeeper), and `drivers/nats` (nats.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`. @@ -98,9 +99,10 @@ Requires JDK 21 (Gradle toolchain auto-downloads if needed). (cd drivers/rabbitmq && go build -o agent .) (cd drivers/rocketmq && go build -o agent .) (cd drivers/zookeeper && go build -o agent .) +(cd drivers/nats && go build -o agent .) ``` -Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/cassandra-go`, `drivers/duckdb`, `drivers/hive-go`, `drivers/iotdb`, `drivers/oracle-go`, `drivers/kingbase-go`, `drivers/vastbase-go`, `drivers/tdengine`, `drivers/xugu`, `drivers/rabbitmq`, `drivers/rocketmq`, and `drivers/zookeeper`. +Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/cassandra-go`, `drivers/duckdb`, `drivers/hive-go`, `drivers/iotdb`, `drivers/oracle-go`, `drivers/kingbase-go`, `drivers/vastbase-go`, `drivers/tdengine`, `drivers/xugu`, `drivers/rabbitmq`, `drivers/rocketmq`, `drivers/zookeeper`, and `drivers/nats`. ### Local DBX Runtime Test @@ -114,7 +116,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 `cassandra`, `hive`, `iotdb`, `oracle`, `kingbase`, `tdengine`, `xugu`, `rabbitmq`, `rocketmq`, and `zookeeper` use an `agent` executable instead of `agent.jar`. TDengine builds `target/release/dbx-tdengine-driver` from `drivers/tdengine/Cargo.toml`. +Native agents such as `cassandra`, `hive`, `iotdb`, `oracle`, `kingbase`, `tdengine`, `xugu`, `rabbitmq`, `rocketmq`, `zookeeper`, and `nats` use an `agent` executable instead of `agent.jar`. TDengine builds `target/release/dbx-tdengine-driver` from `drivers/tdengine/Cargo.toml`. ## Versioning diff --git a/agents/README.zh-CN.md b/agents/README.zh-CN.md index da2f15944a..7cdd7e1119 100644 --- a/agents/README.zh-CN.md +++ b/agents/README.zh-CN.md @@ -46,11 +46,12 @@ DBX 的 Agent 驱动 —— 通过 JDBC 和原生数据库驱动支持各种数 | zookeeper | Apache ZooKeeper | go-zookeeper 原生 Agent | | rabbitmq | RabbitMQ | amqp091-go 原生 agent | | rocketmq | Apache RocketMQ | rocketmq-admin-go 原生 agent | +| nats | NATS | nats.go 原生 agent | ## 多 JRE 支持 -多数 Java agent 以 JRE 21 为目标。原生 agent(如 `cassandra`、`duckdb`、`hive`、`iotdb`、`oracle`、`kingbase`、`tdengine`、`xugu`、`rabbitmq`、`rocketmq` 和 `zookeeper`)不需要 JRE。对 Java agent,DBX 会自动下载并管理 JRE 21 安装。 +多数 Java agent 以 JRE 21 为目标。原生 agent(如 `cassandra`、`duckdb`、`hive`、`iotdb`、`oracle`、`kingbase`、`tdengine`、`xugu`、`rabbitmq`、`rocketmq`、`zookeeper` 和 `nats`)不需要 JRE。对 Java agent,DBX 会自动下载并管理 JRE 21 安装。 ## JDBC 连接池 @@ -76,7 +77,7 @@ HikariCP 会直接打进启用连接池的 Agent JAR。已经使用 DBX 托管 J 对于新 agent,只要存在成熟、许可证兼容的原生驱动,优先选择**原生(Go 或 Rust)驱动**而非 Java/JDBC agent。原生 agent 以单一自包含可执行文件发布,无需 JRE,可显著降低内存占用和启动时间 —— 完全避开 Java agent 即便空闲也要付出的 JVM 基线开销。 -- **原生(Go/Rust)** —— 存在可用原生驱动时首选。参考 `drivers/cassandra-go`(Apache cassandra-gocql-driver)、`drivers/duckdb`、`drivers/hive-go`(原生 HS2)、`drivers/iotdb`(Apache IoTDB Go Client)、`drivers/oracle-go`(go-ora)、`drivers/kingbase-go`(gokb)、`drivers/vastbase-go`(openGauss connector)、`drivers/tdengine`(taos-connector-rust)、`drivers/xugu`、`drivers/rabbitmq`(amqp091-go)、`drivers/rocketmq`(rocketmq-admin-go)和 `drivers/zookeeper`(go-zookeeper)。无需 JRE 下载与管理。 +- **原生(Go/Rust)** —— 存在可用原生驱动时首选。参考 `drivers/cassandra-go`(Apache cassandra-gocql-driver)、`drivers/duckdb`、`drivers/hive-go`(原生 HS2)、`drivers/iotdb`(Apache IoTDB Go Client)、`drivers/oracle-go`(go-ora)、`drivers/kingbase-go`(gokb)、`drivers/vastbase-go`(openGauss connector)、`drivers/tdengine`(taos-connector-rust)、`drivers/xugu`、`drivers/rabbitmq`(amqp091-go)、`drivers/rocketmq`(rocketmq-admin-go)、`drivers/zookeeper`(go-zookeeper)和 `drivers/nats`(nats.go)。无需 JRE 下载与管理。 - **Java/JDBC** —— 当某数据库只有 JDBC 驱动,或原生驱动不成熟、缺乏维护时的默认兜底方案。多数 agent 仍属此类。 原生 agent 实现与 Java agent 相同的 JSON-RPC 契约和 `versions.json` 登记;它发布的是 `agent` 可执行文件而非 `agent.jar`。若同一数据库同时保留原生和 Java 源码实现,默认只发布原生产物;只有 Java 变体以独立兼容配置登记时才同时发布,例如 `oracle-legacy` / `oracle-10g`。 @@ -98,9 +99,10 @@ HikariCP 会直接打进启用连接池的 Agent JAR。已经使用 DBX 托管 J (cd drivers/rabbitmq && go build -o agent .) (cd drivers/rocketmq && go build -o agent .) (cd drivers/zookeeper && go build -o agent .) +(cd drivers/nats && go build -o agent .) ``` -产物 JAR 在 `drivers/{module}/build/libs/`。原生 agent 从 `drivers/cassandra-go`、`drivers/duckdb`、`drivers/hive-go`、`drivers/iotdb`、`drivers/oracle-go`、`drivers/kingbase-go`、`drivers/vastbase-go`、`drivers/tdengine`、`drivers/xugu`、`drivers/rabbitmq`、`drivers/rocketmq` 和 `drivers/zookeeper` 构建。 +产物 JAR 在 `drivers/{module}/build/libs/`。原生 agent 从 `drivers/cassandra-go`、`drivers/duckdb`、`drivers/hive-go`、`drivers/iotdb`、`drivers/oracle-go`、`drivers/kingbase-go`、`drivers/vastbase-go`、`drivers/tdengine`、`drivers/xugu`、`drivers/rabbitmq`、`drivers/rocketmq`、`drivers/zookeeper` 和 `drivers/nats` 构建。 ### 本地 DBX 运行时测试 @@ -114,7 +116,7 @@ cp agents/drivers//build/libs/*-all.jar ~/.dbx/agents/drivers/ 重启 DBX 或断开重连数据库,使新 agent 进程加载替换后的 JAR。 -`cassandra`、`hive`、`iotdb`、`oracle`、`kingbase`、`tdengine`、`xugu`、`rabbitmq`、`rocketmq` 和 `zookeeper` 等原生 agent 使用可执行文件而非 `agent.jar`。TDengine 从 `drivers/tdengine/Cargo.toml` 构建 `target/release/dbx-tdengine-driver`。 +`cassandra`、`hive`、`iotdb`、`oracle`、`kingbase`、`tdengine`、`xugu`、`rabbitmq`、`rocketmq`、`zookeeper` 和 `nats` 等原生 agent 使用可执行文件而非 `agent.jar`。TDengine 从 `drivers/tdengine/Cargo.toml` 构建 `target/release/dbx-tdengine-driver`。 ## 版本管理 diff --git a/agents/drivers/nats/go.mod b/agents/drivers/nats/go.mod new file mode 100644 index 0000000000..0dcb3ae3dc --- /dev/null +++ b/agents/drivers/nats/go.mod @@ -0,0 +1,13 @@ +module github.com/t8y2/dbx/agents/drivers/nats + +go 1.23.0 + +require github.com/nats-io/nats.go v1.47.0 + +require ( + github.com/klauspost/compress v1.18.0 // indirect + github.com/nats-io/nkeys v0.4.11 // indirect + github.com/nats-io/nuid v1.0.1 // indirect + golang.org/x/crypto v0.37.0 // indirect + golang.org/x/sys v0.32.0 // indirect +) diff --git a/agents/drivers/nats/go.sum b/agents/drivers/nats/go.sum new file mode 100644 index 0000000000..76dd281537 --- /dev/null +++ b/agents/drivers/nats/go.sum @@ -0,0 +1,12 @@ +github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= +github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ= +github.com/nats-io/nats.go v1.47.0 h1:YQdADw6J/UfGUd2Oy6tn4Hq6YHxCaJrVKayxxFqYrgM= +github.com/nats-io/nats.go v1.47.0/go.mod h1:iRWIPokVIFbVijxuMQq4y9ttaBTMe0SFdlZfMDd+33g= +github.com/nats-io/nkeys v0.4.11 h1:q44qGV008kYd9W1b1nEBkNzvnWxtRSQ7A8BoqRrcfa0= +github.com/nats-io/nkeys v0.4.11/go.mod h1:szDimtgmfOi9n25JpfIdGw12tZFYXqhGxjhVxsatHVE= +github.com/nats-io/nuid v1.0.1 h1:5iA8DT8V7q8WK2EScv2padNa/rTESc1KdnPw4TC2paw= +github.com/nats-io/nuid v1.0.1/go.mod h1:19wcPz3Ph3q0Jbyiqsd0kePYG7A95tJPxeL+1OSON2c= +golang.org/x/crypto v0.37.0 h1:kJNSjF/Xp7kU0iB2Z+9viTPMW4EqqsrywMXLJOOsXSE= +golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc= +golang.org/x/sys v0.32.0 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20= +golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= diff --git a/agents/drivers/nats/main.go b/agents/drivers/nats/main.go new file mode 100644 index 0000000000..385f058787 --- /dev/null +++ b/agents/drivers/nats/main.go @@ -0,0 +1,1457 @@ +package main + +import ( + "bufio" + "bytes" + "context" + "crypto/tls" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "math" + "net" + "net/http" + "net/url" + "os" + "sort" + "strconv" + "strings" + "sync" + "time" + "unicode" + "unicode/utf8" + + "github.com/nats-io/nats.go" +) + +const ( + maxRPCMessageBytes = 32 * 1024 * 1024 + maxPublishPayloadBytes = 16 * 1024 * 1024 + maxSubjectBytes = 1024 + maxHeaderCount = 100 + maxHeaderKeyBytes = 256 + maxHeaderValueBytes = 8 * 1024 + maxHeaderWireBytes = 64 * 1024 + captureChannelCapacity = 1 + maxCaptureDuration = 60 * time.Second + maxCaptureMessages = 1_000 + maxCaptureBytes = 16 * 1024 * 1024 + maxJetStreamProbeTimeout = time.Second + maxJetStreamListItems = 200 + maxHistoryMessages = 1_000 + maxHistoryBytes = 16 * 1024 * 1024 + maxLivePendingMessages = 1_000 + maxLivePendingBytes = 16 * 1024 * 1024 + maxLiveMessagePayloadBytes = 1 * 1024 * 1024 + maxLiveMessageHeaderBytes = 64 * 1024 +) + +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 captureOptions struct { + subject string + duration time.Duration + maxMessages int + maxBytes int + includeHeaders bool +} + +type historyOptions struct { + stream string + startSequence uint64 + maxMessages int + maxBytes int +} + +type subscriptionInfo struct { + SubscriptionID string `json:"subscriptionId"` + Subject string `json:"subject"` + QueueGroup string `json:"queueGroup,omitempty"` + State string `json:"state"` + ReceivedCount int `json:"receivedCount"` + DroppedCount int `json:"droppedCount"` +} + +type liveSubscription struct { + info subscriptionInfo + nc *nats.Conn + sub *nats.Subscription + seq uint64 +} + +// A server owns only persistent subscriptions. Bounded RPC methods continue +// to create one short-lived connection each, which keeps MCP capture cleanup +// independent from a console subscription lifecycle. +type server struct { + mu sync.Mutex + subscriptions map[string]*liveSubscription + encoder *json.Encoder + writeMu sync.Mutex +} + +func newServer() *server { + return &server{subscriptions: make(map[string]*liveSubscription)} +} + +func main() { + service := newServer() + encoder := json.NewEncoder(os.Stdout) + encoder.SetEscapeHTML(false) + service.encoder = encoder + if _, err := fmt.Fprintln(os.Stdout, `{"ready":true}`); err != nil { + return + } + scanner := bufio.NewScanner(os.Stdin) + scanner.Buffer(make([]byte, 64*1024), maxRPCMessageBytes) + for scanner.Scan() { + // A live subscription emits notifications while ordinary JSON-RPC calls + // must still be able to stop it. Dispatch independently and let clients + // correlate responses by id. + line := append([]byte(nil), scanner.Bytes()...) + go func() { + response, shutdown := service.handle(line) + if err := service.write(response); err != nil { + // This error contains encoder state only; never write request data, + // credentials, or message payloads to stderr. + fmt.Fprintln(os.Stderr, "unable to write NATS agent response") + return + } + if shutdown { + service.close() + } + }() + } + service.close() +} + +func (s *server) handle(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: "invalid JSON-RPC request"} + return response, false + } + if len(request.ID) > 0 { + response.ID = request.ID + } + params := jsonObject{} + if len(request.Params) > 0 && string(request.Params) != "null" { + if err := json.Unmarshal(request.Params, ¶ms); err != nil { + response.Error = &rpcError{Code: -1, Message: "invalid JSON-RPC parameters"} + return response, false + } + } + result, shutdown, err := s.dispatch(request.Method, params) + if err != nil { + response.Error = &rpcError{Code: -1, Message: redactError(err, params)} + return response, false + } + response.Result = result + return response, shutdown +} + +func (s *server) dispatch(method string, params jsonObject) (any, bool, error) { + switch method { + case "handshake": + return handshakeResult{2, 2, []string{ + "test_connection", "nats_core", "nats_headers", "nats_subscription_events", "nats_jetstream_read", + }}, false, nil + case "test_connection": + return s.testConnection(params) + case "publish": + return s.publish(params) + case "capture": + return s.capture(params) + case "start_subscription": + return s.startSubscription(params) + case "stop_subscription": + return s.stopSubscription(params) + case "list_subscriptions": + return s.listSubscriptions(), false, nil + case "jetstream_info": + return s.jetStreamInfo(params) + case "list_streams": + return s.listStreams(params) + case "get_stream": + return s.getStream(params) + case "list_consumers": + return s.listConsumers(params) + case "get_consumer": + return s.getConsumer(params) + case "fetch_history": + return s.fetchHistory(params) + case "shutdown": + return map[string]any{"ok": true}, true, nil + default: + return nil, false, fmt.Errorf("unknown method: %s", method) + } +} + +func (s *server) write(value any) error { + if s.encoder == nil { + return nil + } + s.writeMu.Lock() + defer s.writeMu.Unlock() + return s.encoder.Encode(value) +} + +func (s *server) emit(method string, params any) { + if err := s.write(map[string]any{"jsonrpc": "2.0", "method": method, "params": params}); err != nil { + // Do not include the failing message or request in the log: they can + // contain credentials or application payloads. + fmt.Fprintln(os.Stderr, "unable to write NATS subscription event") + } +} + +func (s *server) startSubscription(params jsonObject) (any, bool, error) { + connection, err := requiredObject(params, "connection") + if err != nil { + return nil, false, err + } + subscription, err := requiredObject(params, "subscription") + if err != nil { + return nil, false, err + } + subscriptionID, err := requiredString(subscription, "subscriptionId") + if err != nil { + return nil, false, err + } + if err := validateSubscriptionID(subscriptionID); err != nil { + return nil, false, err + } + subject, err := requiredString(subscription, "subject") + if err != nil { + return nil, false, err + } + if err := validateSubject(subject, true, "subscription Subject"); err != nil { + return nil, false, err + } + queueGroup, err := optionalString(subscription, "queueGroup") + if err != nil { + return nil, false, err + } + if err := validateQueueGroup(queueGroup); err != nil { + return nil, false, err + } + + s.mu.Lock() + if existing := s.subscriptions[subscriptionID]; existing != nil { + info := existing.info + s.mu.Unlock() + if info.Subject != subject || info.QueueGroup != queueGroup { + return nil, false, errors.New("NATS subscriptionId is already active with a different Subject or queue group") + } + return info, false, nil + } + s.mu.Unlock() + + nc, err := s.connect(connection) + if err != nil { + return nil, false, err + } + live := &liveSubscription{info: subscriptionInfo{SubscriptionID: subscriptionID, Subject: subject, QueueGroup: queueGroup, State: "starting"}, nc: nc} + callback := func(msg *nats.Msg) { s.deliverSubscriptionMessage(subscriptionID, msg) } + if queueGroup == "" { + live.sub, err = nc.Subscribe(subject, callback) + } else { + live.sub, err = nc.QueueSubscribe(subject, queueGroup, callback) + } + if err != nil { + nc.Close() + return nil, false, err + } + if err := live.sub.SetPendingLimits(maxLivePendingMessages, maxLivePendingBytes); err != nil { + live.sub.Unsubscribe() + nc.Close() + return nil, false, err + } + if err := nc.FlushTimeout(requestTimeout(connection)); err != nil { + live.sub.Unsubscribe() + nc.Close() + return nil, false, err + } + + s.mu.Lock() + if existing := s.subscriptions[subscriptionID]; existing != nil { + s.mu.Unlock() + live.sub.Unsubscribe() + nc.Close() + return existing.info, false, nil + } + live.info.State = "active" + s.subscriptions[subscriptionID] = live + info := live.info + live.seq++ + stateSequence := live.seq + // Emit the initial state while holding the subscription lock. A message + // callback cannot overtake it and receive an earlier sequence number. + s.emit("subscription_state", map[string]any{"subscriptionId": subscriptionID, "sequence": stateSequence, "state": "active"}) + s.mu.Unlock() + return info, false, nil +} + +func (s *server) stopSubscription(params jsonObject) (any, bool, error) { + subscriptionID, err := requiredString(params, "subscriptionId") + if err != nil { + return nil, false, err + } + if err := validateSubscriptionID(subscriptionID); err != nil { + return nil, false, err + } + live := s.removeSubscription(subscriptionID) + if live == nil { + return map[string]any{"ok": true}, false, nil + } + if err := live.sub.Unsubscribe(); err != nil && !errors.Is(err, nats.ErrBadSubscription) { + live.nc.Close() + return nil, false, err + } + live.nc.Close() + live.seq++ + s.emit("subscription_state", map[string]any{"subscriptionId": subscriptionID, "sequence": live.seq, "state": "stopped"}) + return map[string]any{"ok": true}, false, nil +} + +func (s *server) listSubscriptions() any { + s.mu.Lock() + defer s.mu.Unlock() + items := make([]subscriptionInfo, 0, len(s.subscriptions)) + for _, live := range s.subscriptions { + info := live.info + if dropped, err := live.sub.Dropped(); err == nil && dropped > info.DroppedCount { + info.DroppedCount = dropped + live.info.DroppedCount = dropped + } + items = append(items, info) + } + sort.Slice(items, func(left, right int) bool { return items[left].SubscriptionID < items[right].SubscriptionID }) + return items +} + +func (s *server) removeSubscription(subscriptionID string) *liveSubscription { + s.mu.Lock() + defer s.mu.Unlock() + live := s.subscriptions[subscriptionID] + delete(s.subscriptions, subscriptionID) + return live +} + +func (s *server) close() { + s.mu.Lock() + subscriptions := s.subscriptions + s.subscriptions = make(map[string]*liveSubscription) + s.mu.Unlock() + for subscriptionID, live := range subscriptions { + _ = live.sub.Unsubscribe() + live.nc.Close() + live.seq++ + s.emit("subscription_state", map[string]any{"subscriptionId": subscriptionID, "sequence": live.seq, "state": "stopped"}) + } +} + +func (s *server) deliverSubscriptionMessage(subscriptionID string, msg *nats.Msg) { + if !liveMessageWithinLimits(msg) { + s.recordLiveMessageDrop(subscriptionID) + return + } + item, _, err := captureMessage(msg, true) + if err != nil { + s.emit("subscription_error", map[string]any{ + "subscriptionId": subscriptionID, + "sequence": s.nextSubscriptionSequence(subscriptionID), + "message": "Unable to decode NATS message", + }) + return + } + s.mu.Lock() + live := s.subscriptions[subscriptionID] + if live == nil { + s.mu.Unlock() + return + } + live.seq++ + live.info.ReceivedCount++ + if dropped, dropErr := live.sub.Dropped(); dropErr == nil && dropped > live.info.DroppedCount { + live.info.DroppedCount = dropped + } + sequence := live.seq + droppedCount := live.info.DroppedCount + s.mu.Unlock() + s.emit("subscription_message", map[string]any{ + "subscriptionId": subscriptionID, + "sequence": sequence, + "droppedCount": droppedCount, + "message": item, + }) +} + +func (s *server) recordLiveMessageDrop(subscriptionID string) { + s.mu.Lock() + defer s.mu.Unlock() + live := s.subscriptions[subscriptionID] + if live == nil { + return + } + live.seq++ + live.info.DroppedCount++ + // Keep the subscription active while surfacing the drop count to the UI. + s.emit("subscription_state", map[string]any{ + "subscriptionId": subscriptionID, + "sequence": live.seq, + "state": "active", + "droppedCount": live.info.DroppedCount, + "detail": "Dropped NATS message that exceeds the live display limit", + }) +} + +func (s *server) nextSubscriptionSequence(subscriptionID string) uint64 { + s.mu.Lock() + defer s.mu.Unlock() + if live := s.subscriptions[subscriptionID]; live != nil { + live.seq++ + return live.seq + } + return 0 +} + +func (s *server) testConnection(params jsonObject) (any, bool, error) { + connection, err := requiredObject(params, "connection") + if err != nil { + return nil, false, err + } + start := time.Now() + nc, err := s.connect(connection) + if err != nil { + return nil, false, err + } + defer nc.Close() + + return map[string]any{ + "ok": true, + "serverName": nc.ConnectedServerName(), + "serverVersion": nc.ConnectedServerVersion(), + "headersSupported": nc.HeadersSupported(), + "jetstreamEnabled": s.probeJetStream(nc, connection), + "maxPayload": nc.MaxPayload(), + "connectedUrl": nc.ConnectedUrl(), + "roundTripMs": time.Since(start).Milliseconds(), + }, false, nil +} + +func (s *server) probeJetStream(nc *nats.Conn, connection jsonObject) bool { + // AccountInfo can wait for a server that has no JetStream responder. Keep + // the capability probe short so a normal connection test remains bounded. + js, err := nc.JetStream(nats.MaxWait(jetStreamProbeTimeout(connection))) + if err != nil { + return false + } + _, err = js.AccountInfo() + return err == nil +} + +// openJetStream keeps all JetStream read RPCs on a short-lived connection. +// Unlike a Core subscription this state is never retained after the response, +// which makes read-only MCP calls independent from the desktop live runtime. +func (s *server) openJetStream(params jsonObject) (nats.JetStreamContext, *nats.Conn, jsonObject, error) { + connection, err := requiredObject(params, "connection") + if err != nil { + return nil, nil, nil, err + } + nc, err := s.connect(connection) + if err != nil { + return nil, nil, nil, err + } + js, err := nc.JetStream(nats.MaxWait(requestTimeout(connection))) + if err != nil { + nc.Close() + return nil, nil, nil, err + } + return js, nc, connection, nil +} + +func (s *server) jetStreamInfo(params jsonObject) (any, bool, error) { + js, nc, connection, err := s.openJetStream(params) + if err != nil { + return nil, false, err + } + defer nc.Close() + ctx, cancel := context.WithTimeout(context.Background(), requestTimeout(connection)) + defer cancel() + info, err := js.AccountInfo(nats.Context(ctx)) + if err != nil { + return nil, false, fmt.Errorf("JetStream account information is unavailable: %w", err) + } + return map[string]any{ + "enabled": true, + "memoryBytes": boundedUint64ToInt64(info.Memory), + "storageBytes": boundedUint64ToInt64(info.Store), + "streams": info.Streams, + "consumers": info.Consumers, + }, false, nil +} + +func (s *server) listStreams(params jsonObject) (any, bool, error) { + js, nc, connection, err := s.openJetStream(params) + if err != nil { + return nil, false, err + } + defer nc.Close() + ctx, cancel := context.WithTimeout(context.Background(), requestTimeout(connection)) + defer cancel() + if err := requireJetStream(js, ctx); err != nil { + return nil, false, err + } + + items := make([]map[string]any, 0) + streamCh := js.Streams(nats.Context(ctx)) + if streamCh == nil { + return nil, false, errors.New("unable to start JetStream stream listing") + } + for stream := range streamCh { + if len(items) == maxJetStreamListItems { + return map[string]any{"streams": items, "truncated": true}, false, nil + } + items = append(items, streamInfoValue(stream)) + } + if err := ctx.Err(); err != nil { + return nil, false, fmt.Errorf("JetStream stream listing timed out: %w", err) + } + return map[string]any{"streams": items, "truncated": false}, false, nil +} + +func (s *server) getStream(params jsonObject) (any, bool, error) { + stream, err := requiredString(params, "stream") + if err != nil { + return nil, false, err + } + if err := validateJetStreamName(stream, "stream"); err != nil { + return nil, false, err + } + js, nc, connection, err := s.openJetStream(params) + if err != nil { + return nil, false, err + } + defer nc.Close() + ctx, cancel := context.WithTimeout(context.Background(), requestTimeout(connection)) + defer cancel() + info, err := js.StreamInfo(stream, nats.Context(ctx)) + if err != nil { + return nil, false, fmt.Errorf("JetStream stream information is unavailable: %w", err) + } + return streamInfoValue(info), false, nil +} + +func (s *server) listConsumers(params jsonObject) (any, bool, error) { + stream, err := requiredString(params, "stream") + if err != nil { + return nil, false, err + } + if err := validateJetStreamName(stream, "stream"); err != nil { + return nil, false, err + } + js, nc, connection, err := s.openJetStream(params) + if err != nil { + return nil, false, err + } + defer nc.Close() + ctx, cancel := context.WithTimeout(context.Background(), requestTimeout(connection)) + defer cancel() + if err := requireJetStream(js, ctx); err != nil { + return nil, false, err + } + + items := make([]map[string]any, 0) + consumerCh := js.Consumers(stream, nats.Context(ctx)) + if consumerCh == nil { + return nil, false, errors.New("unable to start JetStream consumer listing") + } + for consumer := range consumerCh { + if len(items) == maxJetStreamListItems { + return map[string]any{"stream": stream, "consumers": items, "truncated": true}, false, nil + } + items = append(items, consumerInfoValue(consumer)) + } + if err := ctx.Err(); err != nil { + return nil, false, fmt.Errorf("JetStream consumer listing timed out: %w", err) + } + return map[string]any{"stream": stream, "consumers": items, "truncated": false}, false, nil +} + +func (s *server) getConsumer(params jsonObject) (any, bool, error) { + stream, err := requiredString(params, "stream") + if err != nil { + return nil, false, err + } + consumer, err := requiredString(params, "consumer") + if err != nil { + return nil, false, err + } + if err := validateJetStreamName(stream, "stream"); err != nil { + return nil, false, err + } + if err := validateJetStreamName(consumer, "consumer"); err != nil { + return nil, false, err + } + js, nc, connection, err := s.openJetStream(params) + if err != nil { + return nil, false, err + } + defer nc.Close() + ctx, cancel := context.WithTimeout(context.Background(), requestTimeout(connection)) + defer cancel() + info, err := js.ConsumerInfo(stream, consumer, nats.Context(ctx)) + if err != nil { + return nil, false, fmt.Errorf("JetStream consumer information is unavailable: %w", err) + } + return consumerInfoValue(info), false, nil +} + +func (s *server) fetchHistory(params jsonObject) (any, bool, error) { + history, err := requiredObject(params, "history") + if err != nil { + return nil, false, err + } + options, err := parseHistoryOptions(history) + if err != nil { + return nil, false, err + } + js, nc, connection, err := s.openJetStream(params) + if err != nil { + return nil, false, err + } + defer nc.Close() + ctx, cancel := context.WithTimeout(context.Background(), requestTimeout(connection)) + defer cancel() + stream, err := js.StreamInfo(options.stream, nats.Context(ctx)) + if err != nil { + return nil, false, fmt.Errorf("JetStream stream information is unavailable: %w", err) + } + if stream.State.Msgs == 0 || stream.State.FirstSeq == 0 { + return emptyHistoryResult(options.stream), false, nil + } + sequence := options.startSequence + if sequence == 0 || sequence < stream.State.FirstSeq { + sequence = stream.State.FirstSeq + } + messages := make([]map[string]any, 0, options.maxMessages) + bytesUsed := 0 + skipped := 0 + var nextSequence any + truncated := false + for sequence <= stream.State.LastSeq { + message, getErr := js.GetMsg(options.stream, sequence, nats.Context(ctx)) + if getErr != nil { + if errors.Is(getErr, nats.ErrMsgNotFound) { + skipped++ + sequence++ + continue + } + return nil, false, fmt.Errorf("JetStream history read failed: %w", getErr) + } + item, messageBytes, itemErr := streamMessage(message) + if itemErr != nil { + return nil, false, itemErr + } + if messageBytes > options.maxBytes-bytesUsed { + truncated = true + nextSequence = sequence + break + } + messages = append(messages, item) + bytesUsed += messageBytes + sequence++ + if len(messages) == options.maxMessages && sequence <= stream.State.LastSeq { + truncated = true + nextSequence = sequence + break + } + } + return map[string]any{ + "stream": options.stream, + "messages": messages, + "receivedCount": len(messages), + "skippedCount": skipped, + "truncated": truncated, + "nextSequence": nextSequence, + "ackMode": "none", + "consumerKind": "direct_get", + }, false, nil +} + +func emptyHistoryResult(stream string) map[string]any { + return map[string]any{ + "stream": stream, "messages": []map[string]any{}, "receivedCount": 0, "skippedCount": 0, "truncated": false, + "ackMode": "none", "consumerKind": "direct_get", + } +} + +func requireJetStream(js nats.JetStreamContext, ctx context.Context) error { + if _, err := js.AccountInfo(nats.Context(ctx)); err != nil { + return fmt.Errorf("JetStream is unavailable for this connection: %w", err) + } + return nil +} + +func boundedUint64ToInt64(value uint64) int64 { + if value > math.MaxInt64 { + return math.MaxInt64 + } + return int64(value) +} + +func streamInfoValue(info *nats.StreamInfo) map[string]any { + subjects := append([]string(nil), info.Config.Subjects...) + if subjects == nil { + subjects = []string{} + } + return map[string]any{ + "name": info.Config.Name, + "subjects": subjects, + "storage": strings.ToLower(info.Config.Storage.String()), + "retention": strings.ToLower(info.Config.Retention.String()), + "messages": info.State.Msgs, + "bytes": info.State.Bytes, + "firstSequence": info.State.FirstSeq, + "lastSequence": info.State.LastSeq, + "consumers": info.State.Consumers, + } +} + +func consumerInfoValue(info *nats.ConsumerInfo) map[string]any { + return map[string]any{ + "stream": info.Stream, + "name": info.Name, + "filterSubject": info.Config.FilterSubject, + "ackPolicy": strings.TrimPrefix(strings.ToLower(info.Config.AckPolicy.String()), "ack"), + "deliveredConsumerSequence": info.Delivered.Consumer, + "deliveredStreamSequence": info.Delivered.Stream, + "ackFloorConsumerSequence": info.AckFloor.Consumer, + "ackFloorStreamSequence": info.AckFloor.Stream, + "pending": info.NumPending, + "ackPending": info.NumAckPending, + "redelivered": info.NumRedelivered, + } +} + +func (s *server) publish(params jsonObject) (any, bool, error) { + connection, err := requiredObject(params, "connection") + if err != nil { + return nil, false, err + } + publish, err := requiredObject(params, "publish") + if err != nil { + return nil, false, err + } + subject, err := requiredString(publish, "subject") + if err != nil { + return nil, false, err + } + if err := validateSubject(subject, false, "publish Subject"); err != nil { + return nil, false, err + } + reply, err := optionalString(publish, "reply") + if err != nil { + return nil, false, err + } + if reply != "" { + if err := validateSubject(reply, false, "reply Subject"); err != nil { + return nil, false, err + } + } + headers, err := parseHeaders(publish) + if err != nil { + return nil, false, err + } + payloadBase64, err := requiredString(publish, "payloadBase64") + if err != nil { + return nil, false, err + } + payload, err := decodePayloadBase64(payloadBase64) + if err != nil { + return nil, false, err + } + + nc, err := s.connect(connection) + if err != nil { + return nil, false, err + } + defer nc.Close() + if len(headers) > 0 && !nc.HeadersSupported() { + return nil, false, errors.New("NATS server does not support message headers") + } + headerBytes, err := serializedHeaderBytes(headers) + if err != nil { + return nil, false, err + } + if maxPayload := nc.MaxPayload(); maxPayload > 0 && int64(len(payload)+headerBytes) > maxPayload { + return nil, false, fmt.Errorf("NATS message exceeds server maxPayload (%d bytes)", maxPayload) + } + if err := nc.PublishMsg(&nats.Msg{Subject: subject, Reply: reply, Data: payload, Header: headers}); err != nil { + return nil, false, err + } + if err := nc.FlushTimeout(requestTimeout(connection)); err != nil { + return nil, false, err + } + return map[string]any{"acceptedByClient": true, "payloadBytes": len(payload)}, false, nil +} + +func (s *server) capture(params jsonObject) (any, bool, error) { + connection, err := requiredObject(params, "connection") + if err != nil { + return nil, false, err + } + capture, err := requiredObject(params, "capture") + if err != nil { + return nil, false, err + } + options, err := parseCaptureOptions(capture) + if err != nil { + return nil, false, err + } + + nc, err := s.connect(connection) + if err != nil { + return nil, false, err + } + defer nc.Close() + // A one-message channel keeps a high-volume wildcard capture from growing + // memory with maxMessages * serverMaxPayload buffered messages. + ch := make(chan *nats.Msg, captureChannelCapacity) + sub, err := nc.ChanSubscribe(options.subject, ch) + if err != nil { + return nil, false, err + } + defer sub.Unsubscribe() + if err := nc.FlushTimeout(requestTimeout(connection)); err != nil { + return nil, false, err + } + + messages := make([]map[string]any, 0, options.maxMessages) + bytesUsed := 0 + received := 0 + dropped := 0 + timer := time.NewTimer(options.duration) + defer timer.Stop() + stopReason := "duration" + +captureLoop: + for { + select { + case msg := <-ch: + if msg == nil { + continue + } + received++ + item, messageBytes, err := captureMessage(msg, options.includeHeaders) + if err != nil { + return nil, false, err + } + if messageBytes > options.maxBytes-bytesUsed { + dropped++ + stopReason = "byte_limit" + break captureLoop + } + bytesUsed += messageBytes + messages = append(messages, item) + if len(messages) >= options.maxMessages { + stopReason = "message_limit" + break captureLoop + } + if bytesUsed >= options.maxBytes { + stopReason = "byte_limit" + break captureLoop + } + case <-timer.C: + break captureLoop + } + } + + // ChanSubscribe reports overflow through Subscription.Dropped. Read it + // before Unsubscribe closes the subscription and makes the metric invalid. + if clientDropped, dropErr := sub.Dropped(); dropErr == nil && clientDropped > 0 { + dropped += clientDropped + received += clientDropped + } + if err := sub.Unsubscribe(); err != nil { + return nil, false, err + } + return map[string]any{ + "subject": options.subject, + "messages": messages, + "receivedCount": received, + "droppedCount": dropped, + "truncated": stopReason != "duration" || dropped > 0, + "stopReason": stopReason, + }, false, nil +} + +func (s *server) connect(config jsonObject) (*nats.Conn, error) { + endpoint, tlsConfig, err := connectionEndpoint(config) + if err != nil { + return nil, err + } + username, err := optionalString(config, "username") + if err != nil { + return nil, err + } + password, err := optionalString(config, "password") + if err != nil { + return nil, err + } + token, err := optionalString(config, "token") + if err != nil { + return nil, err + } + if token != "" && (username != "" || password != "") { + return nil, errors.New("NATS token authentication cannot be combined with username/password authentication") + } + if username == "" && password != "" { + return nil, errors.New("NATS password authentication requires a username") + } + + opts := []nats.Option{nats.Name("DBX NATS Agent"), nats.Timeout(connectTimeout(config))} + if tlsConfig != nil { + opts = append(opts, nats.Secure(tlsConfig)) + } + if username != "" { + opts = append(opts, nats.UserInfo(username, password)) + } + if token != "" { + opts = append(opts, nats.Token(token)) + } + return nats.Connect(endpoint, opts...) +} + +// connectionEndpoint separates the configured URL (and therefore TLS SNI) +// from the local DBX tunnel endpoint. The latter is never persisted as a +// replacement server URL, which keeps reconnects and certificate validation +// pointed at the logical NATS host. +func connectionEndpoint(config jsonObject) (string, *tls.Config, error) { + serverURL, err := requiredString(config, "serverUrl") + if err != nil { + return "", nil, err + } + parsedURL, err := url.Parse(serverURL) + if err != nil || !matchesNATSScheme(parsedURL.Scheme) || parsedURL.Host == "" { + return "", nil, errors.New("serverUrl must use nats:// or tls:// and include a host") + } + if parsedURL.User != nil { + return "", nil, errors.New("serverUrl must not include credentials; use the dedicated authentication fields") + } + serverName := parsedURL.Hostname() + + connectHost, err := optionalString(config, "connectHost") + if err != nil { + return "", nil, err + } + if connectHost != "" { + connectPort, ok := integerValue(config["connectPort"]) + if !ok || connectPort < 1 || connectPort > 65535 { + return "", nil, errors.New("connectPort must be a valid port when connectHost is configured") + } + parsedURL.Host = net.JoinHostPort(connectHost, strconv.FormatInt(connectPort, 10)) + } else if value, exists := config["connectPort"]; exists && value != nil { + return "", nil, errors.New("connectPort requires connectHost") + } + + usesTLS := parsedURL.Scheme == "tls" + if !usesTLS && boolValue(config, "tlsSkipVerify", false) { + return "", nil, errors.New("tlsSkipVerify requires a tls:// NATS server URL") + } + if !usesTLS { + return parsedURL.String(), nil, nil + } + return parsedURL.String(), &tls.Config{ + ServerName: serverName, + InsecureSkipVerify: boolValue(config, "tlsSkipVerify", false), + }, nil +} + +func parseCaptureOptions(capture jsonObject) (captureOptions, error) { + subject, err := requiredString(capture, "subject") + if err != nil { + return captureOptions{}, err + } + if err := validateSubject(subject, true, "capture Subject"); err != nil { + return captureOptions{}, err + } + durationMs, err := boundedInteger(capture, "durationMs", 5_000, 1, int64(maxCaptureDuration/time.Millisecond)) + if err != nil { + return captureOptions{}, err + } + maxMessages, err := boundedInteger(capture, "maxMessages", 100, 1, maxCaptureMessages) + if err != nil { + return captureOptions{}, err + } + maxBytes, err := boundedInteger(capture, "maxBytes", 1<<20, 1, maxCaptureBytes) + if err != nil { + return captureOptions{}, err + } + includeHeaders, err := optionalBool(capture, "includeHeaders", true) + if err != nil { + return captureOptions{}, err + } + return captureOptions{ + subject: subject, duration: time.Duration(durationMs) * time.Millisecond, + maxMessages: int(maxMessages), maxBytes: int(maxBytes), includeHeaders: includeHeaders, + }, nil +} + +func parseHistoryOptions(history jsonObject) (historyOptions, error) { + stream, err := requiredString(history, "stream") + if err != nil { + return historyOptions{}, err + } + if err := validateJetStreamName(stream, "stream"); err != nil { + return historyOptions{}, err + } + startSequence := uint64(0) + if raw, exists := history["startSequence"]; exists && raw != nil { + value, ok := integerValue(raw) + if !ok || value < 1 { + return historyOptions{}, errors.New("startSequence must be a positive integer when provided") + } + startSequence = uint64(value) + } + maxMessages, err := boundedInteger(history, "maxMessages", 100, 1, maxHistoryMessages) + if err != nil { + return historyOptions{}, err + } + maxBytes, err := boundedInteger(history, "maxBytes", 1<<20, 1, maxHistoryBytes) + if err != nil { + return historyOptions{}, err + } + return historyOptions{ + stream: stream, startSequence: startSequence, maxMessages: int(maxMessages), maxBytes: int(maxBytes), + }, nil +} + +func captureMessage(msg *nats.Msg, includeHeaders bool) (map[string]any, int, error) { + return natsMessageValue(msg.Subject, msg.Reply, msg.Data, msg.Header, time.Now(), includeHeaders) +} + +func liveMessageWithinLimits(msg *nats.Msg) bool { + if len(msg.Data) > maxLiveMessagePayloadBytes { + return false + } + headerBytes, err := serializedHeaderBytes(msg.Header) + return err == nil && headerBytes <= maxLiveMessageHeaderBytes +} + +func streamMessage(msg *nats.RawStreamMsg) (map[string]any, int, error) { + return natsMessageValue(msg.Subject, "", msg.Data, msg.Header, msg.Time, true) +} + +func natsMessageValue( + subject, reply string, data []byte, headers nats.Header, receivedAt time.Time, includeHeaders bool, +) (map[string]any, int, error) { + messageBytes := len(data) + if receivedAt.IsZero() { + receivedAt = time.Now() + } + item := map[string]any{ + "subject": subject, + "payloadBase64": base64.StdEncoding.EncodeToString(data), + "receivedAtMs": receivedAt.UnixMilli(), + "sizeBytes": len(data), + } + if utf8.Valid(data) { + item["payloadText"] = string(data) + } + if includeHeaders && len(headers) > 0 { + headerBytes, err := serializedHeaderBytes(headers) + if err != nil { + return nil, 0, errors.New("received NATS message has invalid headers") + } + messageBytes += headerBytes + item["headers"] = headerArray(headers) + } else if includeHeaders { + item["headers"] = []map[string]string{} + } + if reply != "" { + item["reply"] = reply + } + return item, messageBytes, nil +} + +func validateSubject(subject string, allowWildcards bool, field string) error { + if subject == "" { + return fmt.Errorf("NATS %s is required", field) + } + if len(subject) > maxSubjectBytes { + return fmt.Errorf("NATS %s exceeds the %d byte agent limit", field, maxSubjectBytes) + } + tokens := strings.Split(subject, ".") + for index, token := range tokens { + if token == "" { + return fmt.Errorf("NATS %s cannot contain empty tokens", field) + } + switch token { + case "*": + if allowWildcards { + continue + } + return fmt.Errorf("NATS %s must be a concrete subject", field) + case ">": + if allowWildcards && index == len(tokens)-1 { + continue + } + return fmt.Errorf("NATS %s wildcard > must be the final token", field) + } + if strings.ContainsAny(token, "*>") { + return fmt.Errorf("NATS %s has an invalid wildcard token", field) + } + for _, character := range token { + if unicode.IsSpace(character) || unicode.IsControl(character) { + return fmt.Errorf("NATS %s cannot contain whitespace or control characters", field) + } + } + } + return nil +} + +func validateJetStreamName(value, kind string) error { + if value == "" || len(value) > 256 { + return fmt.Errorf("NATS JetStream %s name is required and must be at most 256 bytes", kind) + } + for _, character := range value { + if unicode.IsSpace(character) || unicode.IsControl(character) || strings.ContainsRune(".*>/\\", character) { + return fmt.Errorf("NATS JetStream %s name contains unsupported characters", kind) + } + } + return nil +} + +func validateSubscriptionID(value string) error { + if value == "" || len(value) > 128 { + return errors.New("NATS subscriptionId is required and must be at most 128 bytes") + } + for _, character := range value { + if unicode.IsSpace(character) || unicode.IsControl(character) { + return errors.New("NATS subscriptionId cannot contain whitespace or control characters") + } + } + return nil +} + +func validateQueueGroup(value string) error { + if value == "" { + return nil + } + if len(value) > 256 { + return errors.New("NATS queue group exceeds the 256 byte agent limit") + } + for _, character := range value { + if unicode.IsSpace(character) || unicode.IsControl(character) { + return errors.New("NATS queue group cannot contain whitespace or control characters") + } + } + return nil +} + +func parseHeaders(parent jsonObject) (nats.Header, error) { + raw, exists := parent["headers"] + if !exists || raw == nil { + return nil, nil + } + values, ok := raw.([]any) + if !ok { + return nil, errors.New("NATS headers must be an array") + } + if len(values) > maxHeaderCount { + return nil, fmt.Errorf("NATS headers cannot contain more than %d entries", maxHeaderCount) + } + result := nats.Header{} + for _, rawValue := range values { + item, ok := asObject(rawValue) + if !ok { + return nil, errors.New("each NATS header must be an object") + } + key, err := requiredString(item, "key") + if err != nil { + return nil, fmt.Errorf("invalid NATS header key: %w", err) + } + value, err := requiredString(item, "value") + if err != nil { + return nil, fmt.Errorf("invalid NATS header value: %w", err) + } + if !validHeaderKey(key) || len(key) > maxHeaderKeyBytes { + return nil, errors.New("NATS header key must be a valid, bounded HTTP token") + } + if !validHeaderValue(value) || len(value) > maxHeaderValueBytes { + return nil, errors.New("NATS header value contains invalid or oversized data") + } + result.Add(key, value) + } + headerBytes, err := serializedHeaderBytes(result) + if err != nil { + return nil, errors.New("NATS headers are invalid") + } + if headerBytes > maxHeaderWireBytes { + return nil, fmt.Errorf("NATS headers exceed the %d byte agent limit", maxHeaderWireBytes) + } + return result, nil +} + +func validHeaderKey(key string) bool { + if key == "" { + return false + } + for _, character := range key { + if !(character >= '0' && character <= '9' || character >= 'A' && character <= 'Z' || character >= 'a' && character <= 'z' || strings.ContainsRune("!#$%&'*+-.^_`|~", character)) { + return false + } + } + return true +} + +func validHeaderValue(value string) bool { + for _, character := range value { + if character == '\r' || character == '\n' || character == 0 || character == 0x7f || (character < 0x20 && character != '\t') { + return false + } + } + return true +} + +func serializedHeaderBytes(headers nats.Header) (int, error) { + if len(headers) == 0 { + return 0, nil + } + var buffer bytes.Buffer + if _, err := buffer.WriteString("NATS/1.0\r\n"); err != nil { + return 0, err + } + if err := http.Header(headers).Write(&buffer); err != nil { + return 0, err + } + if _, err := buffer.WriteString("\r\n"); err != nil { + return 0, err + } + return buffer.Len(), nil +} + +func decodePayloadBase64(encoded string) ([]byte, error) { + if len(encoded) > base64.StdEncoding.EncodedLen(maxPublishPayloadBytes) { + return nil, fmt.Errorf("NATS payload exceeds the %d byte agent limit", maxPublishPayloadBytes) + } + if strings.ContainsAny(encoded, "\r\n") { + return nil, errors.New("payloadBase64 must be canonical base64") + } + payload, err := base64.StdEncoding.Strict().DecodeString(encoded) + if err != nil || base64.StdEncoding.EncodeToString(payload) != encoded { + return nil, errors.New("payloadBase64 must be canonical base64") + } + if len(payload) > maxPublishPayloadBytes { + return nil, fmt.Errorf("NATS payload exceeds the %d byte agent limit", maxPublishPayloadBytes) + } + return payload, nil +} + +func requiredObject(parent jsonObject, key string) (jsonObject, error) { + value, exists := parent[key] + if !exists || value == nil { + return nil, fmt.Errorf("%s is required", key) + } + if object, ok := asObject(value); ok { + return object, nil + } + return nil, fmt.Errorf("%s must be an object", key) +} + +func asObject(value any) (jsonObject, bool) { + switch object := value.(type) { + case jsonObject: + return object, true + case map[string]any: + return jsonObject(object), true + default: + return nil, false + } +} + +func requiredString(object jsonObject, key string) (string, error) { + value, exists := object[key] + if !exists || value == nil { + return "", fmt.Errorf("%s is required", key) + } + stringValue, ok := value.(string) + if !ok { + return "", fmt.Errorf("%s must be a string", key) + } + return stringValue, nil +} + +func optionalString(object jsonObject, key string) (string, error) { + value, exists := object[key] + if !exists || value == nil { + return "", nil + } + stringValue, ok := value.(string) + if !ok { + return "", fmt.Errorf("%s must be a string", key) + } + return stringValue, nil +} + +func optionalBool(object jsonObject, key string, fallback bool) (bool, error) { + value, exists := object[key] + if !exists || value == nil { + return fallback, nil + } + boolean, ok := value.(bool) + if !ok { + return false, fmt.Errorf("%s must be a boolean", key) + } + return boolean, nil +} + +func boolValue(object jsonObject, key string, fallback bool) bool { + value, err := optionalBool(object, key, fallback) + if err != nil { + return fallback + } + return value +} + +func boundedInteger(object jsonObject, key string, fallback, minimum, maximum int64) (int64, error) { + value, exists := object[key] + if !exists || value == nil { + return fallback, nil + } + integer, ok := integerValue(value) + if !ok { + return 0, fmt.Errorf("%s must be an integer", key) + } + if integer < minimum || integer > maximum { + return 0, fmt.Errorf("%s must be between %d and %d", key, minimum, maximum) + } + return integer, nil +} + +func integerValue(value any) (int64, bool) { + switch number := value.(type) { + case int: + return int64(number), true + case int8: + return int64(number), true + case int16: + return int64(number), true + case int32: + return int64(number), true + case int64: + return number, true + case uint: + if uint64(number) <= math.MaxInt64 { + return int64(number), true + } + case uint8: + return int64(number), true + case uint16: + return int64(number), true + case uint32: + return int64(number), true + case uint64: + if number <= math.MaxInt64 { + return int64(number), true + } + case float64: + if !math.IsNaN(number) && !math.IsInf(number, 0) && math.Trunc(number) == number && number >= math.MinInt64 && number <= math.MaxInt64 { + return int64(number), true + } + case json.Number: + integer, err := number.Int64() + if err == nil { + return integer, true + } + } + return 0, false +} + +func requestTimeout(config jsonObject) time.Duration { + ms, ok := integerValue(config["requestTimeoutMs"]) + if !ok || ms <= 0 { + ms = 30_000 + } + return time.Duration(ms) * time.Millisecond +} + +func connectTimeout(config jsonObject) time.Duration { + ms, ok := integerValue(config["connectTimeoutMs"]) + if !ok || ms <= 0 { + ms = 15_000 + } + return time.Duration(ms) * time.Millisecond +} + +func jetStreamProbeTimeout(config jsonObject) time.Duration { + timeout := requestTimeout(config) + if timeout > maxJetStreamProbeTimeout { + return maxJetStreamProbeTimeout + } + return timeout +} + +func matchesNATSScheme(scheme string) bool { + return scheme == "nats" || scheme == "tls" +} + +func headerArray(header nats.Header) []map[string]string { + keys := make([]string, 0, len(header)) + for key := range header { + keys = append(keys, key) + } + sort.Strings(keys) + result := make([]map[string]string, 0) + for _, key := range keys { + for _, value := range header[key] { + result = append(result, map[string]string{"key": key, "value": value}) + } + } + return result +} + +func redactError(err error, params jsonObject) string { + message := err.Error() + connection, ok := asObject(params["connection"]) + if !ok { + return message + } + for _, key := range []string{"password", "token"} { + if value, valueErr := optionalString(connection, key); valueErr == nil && value != "" { + message = strings.ReplaceAll(message, value, "[REDACTED]") + } + } + if serverURL, urlErr := optionalString(connection, "serverUrl"); urlErr == nil { + if parsedURL, parseErr := url.Parse(serverURL); parseErr == nil && parsedURL.User != nil { + message = strings.ReplaceAll(message, parsedURL.User.String(), "[REDACTED]") + } + } + return message +} diff --git a/agents/drivers/nats/main_test.go b/agents/drivers/nats/main_test.go new file mode 100644 index 0000000000..210ba14bdf --- /dev/null +++ b/agents/drivers/nats/main_test.go @@ -0,0 +1,607 @@ +package main + +import ( + "bufio" + "bytes" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "strconv" + "strings" + "sync" + "testing" + "time" + + "github.com/nats-io/nats.go" +) + +func TestHandshakeDispatch(t *testing.T) { + service := &server{} + response, shutdown := service.handle([]byte(`{"jsonrpc":"2.0","id":7,"method":"handshake","params":{"client":"dbx"}}`)) + if shutdown { + t.Fatal("handshake must not shut down the agent") + } + if response.Error != nil { + t.Fatalf("handshake returned error: %+v", response.Error) + } + var result handshakeResult + encoded, err := json.Marshal(response.Result) + if err != nil { + t.Fatal(err) + } + if err := json.Unmarshal(encoded, &result); err != nil { + t.Fatal(err) + } + if result.ProtocolVersion != 2 || result.AgentProtocolVersion != 2 { + t.Fatalf("unexpected handshake versions: %+v", result) + } + if !contains(result.Capabilities, "nats_jetstream_read") { + t.Fatalf("JetStream read capability must be advertised: %+v", result.Capabilities) + } +} + +func TestJetStreamNamesAndHistoryLimitsAreValidatedBeforeConnecting(t *testing.T) { + for _, name := range []string{"ORDERS", "orders-archive", "tenant_orders"} { + if err := validateJetStreamName(name, "stream"); err != nil { + t.Fatalf("JetStream name %q should be valid: %v", name, err) + } + } + for _, name := range []string{"", "orders.stream", "orders/stream", "orders stream", "orders.>"} { + if err := validateJetStreamName(name, "stream"); err == nil { + t.Fatalf("JetStream name %q should be rejected", name) + } + } + for _, history := range []jsonObject{ + {"stream": "ORDERS", "maxMessages": 1_001}, + {"stream": "ORDERS", "maxBytes": maxHistoryBytes + 1}, + {"stream": "ORDERS", "startSequence": 0}, + } { + if _, err := parseHistoryOptions(history); err == nil { + t.Fatalf("history options %#v should be rejected", history) + } + } + + service := &server{} + _, _, err := service.dispatch("fetch_history", jsonObject{ + "connection": jsonObject{"serverUrl": "nats://127.0.0.1:1"}, + "history": jsonObject{"stream": "orders.stream"}, + }) + if err == nil { + t.Fatal("invalid history stream must be rejected before connecting") + } +} + +func TestJetStreamValuesAndHistoryMessageUseSharedDTOShape(t *testing.T) { + stream := streamInfoValue(&nats.StreamInfo{ + Config: nats.StreamConfig{Name: "ORDERS", Subjects: []string{"orders.created"}, Storage: nats.FileStorage, Retention: nats.LimitsPolicy}, + State: nats.StreamState{Msgs: 3, Bytes: 42, FirstSeq: 4, LastSeq: 6, Consumers: 2}, + }) + if stream["storage"] != "file" || stream["retention"] != "limits" || stream["firstSequence"] != uint64(4) { + t.Fatalf("unexpected stream DTO: %#v", stream) + } + consumer := consumerInfoValue(&nats.ConsumerInfo{ + Stream: "ORDERS", Name: "DASHBOARD", Config: nats.ConsumerConfig{FilterSubject: "orders.>", AckPolicy: nats.AckExplicitPolicy}, + Delivered: nats.SequenceInfo{Consumer: 8, Stream: 10}, AckFloor: nats.SequenceInfo{Consumer: 6, Stream: 8}, + NumPending: 4, NumAckPending: 2, NumRedelivered: 1, + }) + if consumer["ackPolicy"] != "explicit" || consumer["pending"] != uint64(4) || consumer["ackPending"] != 2 { + t.Fatalf("unexpected consumer DTO: %#v", consumer) + } + message, messageBytes, err := streamMessage(&nats.RawStreamMsg{ + Subject: "orders.created", Sequence: 4, Header: nats.Header{"Nats-Msg-Id": []string{"42"}}, Data: []byte("ok"), + Time: time.UnixMilli(1_700_000_000_000), + }) + if err != nil || message["payloadBase64"] != "b2s=" || message["receivedAtMs"] != int64(1_700_000_000_000) || messageBytes <= 2 { + t.Fatalf("unexpected history message: value=%#v bytes=%d err=%v", message, messageBytes, err) + } +} + +func TestJetStreamReadRPCsUseOnlyReadAPIs(t *testing.T) { + service := &server{} + connection := jsonObject{ + "serverUrl": startFakeNATSServer(t, fakeNATSOptions{headersSupported: true, jetStream: true}), "requestTimeoutMs": 500, + } + info, _, err := service.dispatch("jetstream_info", jsonObject{"connection": connection}) + if err != nil || !info.(map[string]any)["enabled"].(bool) { + t.Fatalf("JetStream info must succeed: value=%#v err=%v", info, err) + } + streams, _, err := service.dispatch("list_streams", jsonObject{"connection": connection}) + if err != nil || len(streams.(map[string]any)["streams"].([]map[string]any)) != 1 { + t.Fatalf("JetStream Stream list must use API response: value=%#v err=%v", streams, err) + } + stream, _, err := service.dispatch("get_stream", jsonObject{"connection": connection, "stream": "ORDERS"}) + if err != nil || stream.(map[string]any)["lastSequence"] != uint64(2) { + t.Fatalf("JetStream Stream info must succeed: value=%#v err=%v", stream, err) + } + consumers, _, err := service.dispatch("list_consumers", jsonObject{"connection": connection, "stream": "ORDERS"}) + if err != nil || len(consumers.(map[string]any)["consumers"].([]map[string]any)) != 1 { + t.Fatalf("JetStream Consumer list must use API response: value=%#v err=%v", consumers, err) + } + consumer, _, err := service.dispatch("get_consumer", jsonObject{ + "connection": connection, "stream": "ORDERS", "consumer": "DASHBOARD", + }) + if err != nil || consumer.(map[string]any)["ackPolicy"] != "explicit" { + t.Fatalf("JetStream Consumer info must succeed: value=%#v err=%v", consumer, err) + } + history, _, err := service.dispatch("fetch_history", jsonObject{ + "connection": connection, "history": jsonObject{"stream": "ORDERS", "maxMessages": 1, "maxBytes": 1024}, + }) + if err != nil { + t.Fatalf("JetStream direct history read must succeed: %v", err) + } + historyResult := history.(map[string]any) + if historyResult["ackMode"] != "none" || historyResult["consumerKind"] != "direct_get" || !historyResult["truncated"].(bool) { + t.Fatalf("history must be bounded and side-effect free: %#v", historyResult) + } +} + +func contains(values []string, value string) bool { + for _, item := range values { + if item == value { + return true + } + } + return false +} + +func TestPublishRejectsWildcardBeforeConnecting(t *testing.T) { + service := &server{} + _, _, err := service.dispatch("publish", jsonObject{ + "connection": jsonObject{"serverUrl": "nats://127.0.0.1:1"}, + "publish": jsonObject{"subject": "orders.>", "payloadBase64": ""}, + }) + if err == nil { + t.Fatal("wildcard publish should fail before attempting a connection") + } +} + +func TestCaptureRejectsUnboundedLimitsBeforeConnecting(t *testing.T) { + service := &server{} + _, _, err := service.dispatch("capture", jsonObject{ + "connection": jsonObject{"serverUrl": "nats://127.0.0.1:1"}, + "capture": jsonObject{"subject": "orders.>", "durationMs": 60_001}, + }) + if err == nil { + t.Fatal("capture with an excessive duration should be rejected") + } +} + +func TestLiveMessagesAreBoundedBeforeTheyReachTheEventStream(t *testing.T) { + if !liveMessageWithinLimits(&nats.Msg{Data: make([]byte, maxLiveMessagePayloadBytes)}) { + t.Fatal("a live message at the payload limit must be accepted") + } + if liveMessageWithinLimits(&nats.Msg{Data: make([]byte, maxLiveMessagePayloadBytes+1)}) { + t.Fatal("an oversized live payload must be dropped before event encoding") + } + largeHeader := nats.Header{"X-Large": []string{strings.Repeat("a", maxLiveMessageHeaderBytes)}} + if liveMessageWithinLimits(&nats.Msg{Header: largeHeader}) { + t.Fatal("an oversized live header block must be dropped before event encoding") + } +} + +func TestMalformedRPCReturnsError(t *testing.T) { + service := &server{} + response, shutdown := service.handle([]byte("not-json")) + if shutdown || response.Error == nil { + t.Fatalf("malformed RPC should return an error response: %+v", response) + } +} + +func TestValidateSubjectEnforcesNATSTokens(t *testing.T) { + for _, subject := range []string{"orders.created", "orders.*", "orders.*.created", "orders.>", ">"} { + if err := validateSubject(subject, true, "capture Subject"); err != nil { + t.Fatalf("capture subject %q should be valid: %v", subject, err) + } + } + for _, subject := range []string{"", "orders..created", "orders.>.created", "orders.foo*", "orders. created", "orders.\ncreated"} { + if err := validateSubject(subject, true, "capture Subject"); err == nil { + t.Fatalf("capture subject %q should be rejected", subject) + } + } + for _, subject := range []string{"orders.*", "orders.>"} { + if err := validateSubject(subject, false, "publish Subject"); err == nil { + t.Fatalf("publish subject %q should reject wildcards", subject) + } + } +} + +func TestHeaderAndPayloadValidation(t *testing.T) { + headers, err := parseHeaders(jsonObject{"headers": []any{ + jsonObject{"key": "Nats-Msg-Id", "value": "42"}, + jsonObject{"key": "Nats-Msg-Id", "value": "43"}, + }}) + if err != nil || len(headers["Nats-Msg-Id"]) != 2 { + t.Fatalf("valid repeated headers should be preserved: headers=%#v err=%v", headers, err) + } + for _, header := range []jsonObject{ + {"key": "", "value": "value"}, + {"key": "bad key", "value": "value"}, + {"key": "X-Test", "value": "line\r\nbreak"}, + } { + if _, err := parseHeaders(jsonObject{"headers": []any{header}}); err == nil { + t.Fatalf("header %#v should be rejected", header) + } + } + + if _, err := decodePayloadBase64("not base64"); err == nil { + t.Fatal("malformed base64 should be rejected") + } + if _, err := decodePayloadBase64("YWJj\n"); err == nil { + t.Fatal("non-canonical base64 should be rejected") + } + overLimit := strings.Repeat("A", base64.StdEncoding.EncodedLen(maxPublishPayloadBytes)+1) + if _, err := decodePayloadBase64(overLimit); err == nil { + t.Fatal("payload over the Agent limit should be rejected") + } +} + +func TestPublishRejectsHeadersWhenServerDoesNotSupportThem(t *testing.T) { + service := &server{} + _, _, err := service.dispatch("publish", jsonObject{ + "connection": jsonObject{"serverUrl": startFakeNATSServer(t, fakeNATSOptions{headersSupported: false})}, + "publish": jsonObject{ + "subject": "orders.created", "payloadBase64": "", + "headers": []any{jsonObject{"key": "Nats-Msg-Id", "value": "42"}}, + }, + }) + if err == nil || !strings.Contains(err.Error(), "does not support message headers") { + t.Fatalf("expected explicit headers-unsupported error, got %v", err) + } +} + +func TestPublishRejectsPayloadOverServerMaxPayload(t *testing.T) { + service := &server{} + _, _, err := service.dispatch("publish", jsonObject{ + "connection": jsonObject{"serverUrl": startFakeNATSServer(t, fakeNATSOptions{headersSupported: true, maxPayload: 3})}, + "publish": jsonObject{"subject": "orders.created", "payloadBase64": base64.StdEncoding.EncodeToString([]byte("four"))}, + }) + if err == nil || !strings.Contains(err.Error(), "exceeds server maxPayload") { + t.Fatalf("expected server maxPayload error, got %v", err) + } +} + +func TestCaptureStopsAtMessageAndByteLimits(t *testing.T) { + service := &server{} + messageLimitedURL := startFakeNATSServer(t, fakeNATSOptions{messages: [][]byte{[]byte("one"), []byte("two")}}) + result, _, err := service.dispatch("capture", jsonObject{ + "connection": jsonObject{"serverUrl": messageLimitedURL, "requestTimeoutMs": 500}, + "capture": jsonObject{"subject": "orders.created", "durationMs": 500, "maxMessages": 2, "maxBytes": 32}, + }) + if err != nil { + t.Fatalf("message-limited capture failed: %v", err) + } + messageLimited := result.(map[string]any) + if messageLimited["stopReason"] != "message_limit" || !messageLimited["truncated"].(bool) || len(messageLimited["messages"].([]map[string]any)) != 2 { + t.Fatalf("unexpected message-limited capture result: %#v", messageLimited) + } + + byteLimitedURL := startFakeNATSServer(t, fakeNATSOptions{messages: [][]byte{[]byte("five!")}}) + result, _, err = service.dispatch("capture", jsonObject{ + "connection": jsonObject{"serverUrl": byteLimitedURL, "requestTimeoutMs": 500}, + "capture": jsonObject{"subject": "orders.created", "durationMs": 500, "maxMessages": 2, "maxBytes": 4}, + }) + if err != nil { + t.Fatalf("byte-limited capture failed: %v", err) + } + byteLimited := result.(map[string]any) + if byteLimited["stopReason"] != "byte_limit" || !byteLimited["truncated"].(bool) || byteLimited["droppedCount"].(int) != 1 { + t.Fatalf("unexpected byte-limited capture result: %#v", byteLimited) + } +} + +func TestJetStreamProbeUsesRequestBound(t *testing.T) { + if timeout := jetStreamProbeTimeout(jsonObject{"requestTimeoutMs": 30_000}); timeout != maxJetStreamProbeTimeout { + t.Fatalf("JetStream probe timeout should be capped at %s, got %s", maxJetStreamProbeTimeout, timeout) + } + service := &server{} + start := time.Now() + result, _, err := service.dispatch("test_connection", jsonObject{ + "connection": jsonObject{ + "serverUrl": startFakeNATSServer(t, fakeNATSOptions{headersSupported: true}), + "requestTimeoutMs": 75, + }, + }) + if err != nil { + t.Fatalf("test connection failed: %v", err) + } + if elapsed := time.Since(start); elapsed > 750*time.Millisecond { + t.Fatalf("JetStream probe exceeded its request bound: %s", elapsed) + } + if result.(map[string]any)["jetstreamEnabled"].(bool) { + t.Fatalf("server without a JetStream response must not be reported as enabled: %#v", result) + } +} + +func TestRedactErrorDoesNotExposeCredentials(t *testing.T) { + params := jsonObject{"connection": jsonObject{ + "serverUrl": "nats://alice:hunter2@localhost:4222", + "password": "hunter2", + "token": "token-secret", + }} + message := redactError(errors.New("failed nats://alice:hunter2@localhost:4222 with hunter2 and token-secret"), params) + if strings.Contains(message, "hunter2") || strings.Contains(message, "token-secret") { + t.Fatalf("credential leaked in error message: %q", message) + } +} + +func TestConnectionEndpointKeepsTLSNameWhenUsingTunnel(t *testing.T) { + endpoint, tlsConfig, err := connectionEndpoint(jsonObject{ + "serverUrl": "tls://nats.example.test:4222", + "connectHost": "127.0.0.1", + "connectPort": 43123, + "tlsSkipVerify": false, + }) + if err != nil { + t.Fatalf("tunnel endpoint should be valid: %v", err) + } + if endpoint != "tls://127.0.0.1:43123" { + t.Fatalf("unexpected tunnel endpoint: %q", endpoint) + } + if tlsConfig == nil || tlsConfig.ServerName != "nats.example.test" || tlsConfig.InsecureSkipVerify { + t.Fatalf("TLS must preserve the configured server name: %#v", tlsConfig) + } +} + +func TestConnectionEndpointRejectsIncompleteOrInsecureOverrides(t *testing.T) { + for _, config := range []jsonObject{ + {"serverUrl": "nats://localhost:4222", "connectPort": 4223}, + {"serverUrl": "nats://localhost:4222", "connectHost": "127.0.0.1", "connectPort": 0}, + {"serverUrl": "nats://localhost:4222", "tlsSkipVerify": true}, + } { + if _, _, err := connectionEndpoint(config); err == nil { + t.Fatalf("connection override %#v should be rejected", config) + } + } +} + +func TestPersistentSubscriptionLifecycleIsIdempotent(t *testing.T) { + service := newServer() + params := jsonObject{ + "connection": jsonObject{"serverUrl": startFakeNATSServer(t, fakeNATSOptions{headersSupported: true})}, + "subscription": jsonObject{"subscriptionId": "sub-1", "subject": "orders.>"}, + } + result, _, err := service.dispatch("start_subscription", params) + if err != nil { + t.Fatalf("start subscription failed: %v", err) + } + if info := result.(subscriptionInfo); info.State != "active" || info.SubscriptionID != "sub-1" { + t.Fatalf("unexpected subscription info: %#v", info) + } + result, _, err = service.dispatch("start_subscription", params) + if err != nil { + t.Fatalf("duplicate subscription start failed: %v", err) + } + if info := result.(subscriptionInfo); info.SubscriptionID != "sub-1" { + t.Fatalf("duplicate start should return the existing subscription: %#v", info) + } + result, _, err = service.dispatch("list_subscriptions", jsonObject{}) + if err != nil { + t.Fatalf("list subscriptions failed: %v", err) + } + if items := result.([]subscriptionInfo); len(items) != 1 || items[0].SubscriptionID != "sub-1" { + t.Fatalf("unexpected active subscriptions: %#v", items) + } + result, _, err = service.dispatch("stop_subscription", jsonObject{"subscriptionId": "sub-1"}) + if err != nil || !result.(map[string]any)["ok"].(bool) { + t.Fatalf("stop subscription failed: result=%#v error=%v", result, err) + } + result, _, err = service.dispatch("list_subscriptions", jsonObject{}) + if err != nil || len(result.([]subscriptionInfo)) != 0 { + t.Fatalf("subscription must be removed after stop: result=%#v error=%v", result, err) + } +} + +func TestPersistentSubscriptionRejectsInvalidIdsAndQueueGroups(t *testing.T) { + service := newServer() + base := jsonObject{ + "connection": jsonObject{"serverUrl": "nats://127.0.0.1:1"}, + "subscription": jsonObject{"subscriptionId": "sub-1", "subject": "orders.created"}, + } + for _, subscription := range []jsonObject{ + {"subscriptionId": "bad id", "subject": "orders.created"}, + {"subscriptionId": "sub-1", "subject": "orders.>.created"}, + {"subscriptionId": "sub-1", "subject": "orders.created", "queueGroup": "queue group"}, + } { + params := jsonObject{"connection": base["connection"], "subscription": subscription} + if _, _, err := service.dispatch("start_subscription", params); err == nil { + t.Fatalf("invalid subscription %#v must be rejected before connecting", subscription) + } + } +} + +type fakeNATSOptions struct { + headersSupported bool + maxPayload int64 + messages [][]byte + jetStream bool +} + +func startFakeNATSServer(t *testing.T, options fakeNATSOptions) string { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + if options.maxPayload == 0 { + options.maxPayload = 1024 * 1024 + } + done := make(chan struct{}) + var connections sync.WaitGroup + go func() { + defer close(done) + for { + connection, err := listener.Accept() + if err != nil { + return + } + connections.Add(1) + go func() { + defer connections.Done() + serveFakeNATSConnection(connection, options) + }() + } + }() + t.Cleanup(func() { + _ = listener.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Error("fake NATS server did not stop") + } + waitForConnections := make(chan struct{}) + go func() { + connections.Wait() + close(waitForConnections) + }() + select { + case <-waitForConnections: + case <-time.After(time.Second): + t.Error("fake NATS connections did not stop") + } + }) + return "nats://" + listener.Addr().String() +} + +func serveFakeNATSConnection(connection net.Conn, options fakeNATSOptions) { + defer connection.Close() + info, _ := json.Marshal(map[string]any{ + "server_id": "DBX-TEST", "server_name": "dbx-test", "version": "2.10.0", "proto": 1, + "headers": options.headersSupported, "max_payload": options.maxPayload, + }) + if _, err := fmt.Fprintf(connection, "INFO %s\r\n", info); err != nil { + return + } + reader := bufio.NewReader(connection) + var subscriptionSubject, subscriptionID string + subscriptions := make(map[string]string) + messagesSent := false + for { + line, err := reader.ReadString('\n') + if err != nil { + return + } + fields := strings.Fields(line) + if len(fields) == 0 { + continue + } + switch fields[0] { + case "PING": + if _, err := fmt.Fprint(connection, "PONG\r\n"); err != nil { + return + } + if subscriptionID != "" && !messagesSent { + messagesSent = true + for _, payload := range options.messages { + time.Sleep(8 * time.Millisecond) + if _, err := fmt.Fprintf(connection, "MSG %s %s %d\r\n%s\r\n", subscriptionSubject, subscriptionID, len(payload), payload); err != nil { + return + } + } + } + case "SUB": + if len(fields) >= 3 { + subject := fields[len(fields)-2] + sid := fields[len(fields)-1] + subscriptions[subject] = sid + if !strings.HasPrefix(subject, "_INBOX.") { + subscriptionSubject = subject + subscriptionID = sid + } + } + case "PUB", "HPUB": + payload, err := readNATSPayload(reader, fields) + if err != nil { + return + } + if options.jetStream && fields[0] == "PUB" && len(fields) == 4 && strings.HasPrefix(fields[1], "$JS.API.") { + reply := fields[2] + if sid := fakeSubscriptionID(subscriptions, reply); sid != "" { + response, responseErr := fakeJetStreamResponse(fields[1], payload) + if responseErr != nil || writeFakeMessage(connection, reply, sid, response) != nil { + return + } + } + } + } + } +} + +func fakeSubscriptionID(subscriptions map[string]string, reply string) string { + if sid := subscriptions[reply]; sid != "" { + return sid + } + for subject, sid := range subscriptions { + if strings.HasSuffix(subject, ".*") && strings.HasPrefix(reply, strings.TrimSuffix(subject, "*")) { + return sid + } + } + return "" +} + +func fakeJetStreamResponse(subject string, request []byte) ([]byte, error) { + stream := map[string]any{ + "config": map[string]any{"name": "ORDERS", "subjects": []string{"orders.created"}, "storage": "file", "retention": "limits"}, + "state": map[string]any{"messages": 2, "bytes": 4, "first_seq": 1, "last_seq": 2, "consumer_count": 1}, + } + consumer := map[string]any{ + "stream_name": "ORDERS", "name": "DASHBOARD", + "config": map[string]any{"filter_subject": "orders.>", "ack_policy": "explicit"}, + "delivered": map[string]any{"consumer_seq": 2, "stream_seq": 2}, + "ack_floor": map[string]any{"consumer_seq": 1, "stream_seq": 1}, + "num_pending": 1, "num_ack_pending": 1, "num_redelivered": 0, + } + var response any + switch subject { + case "$JS.API.INFO": + response = map[string]any{"memory": 1, "storage": 4, "streams": 1, "consumers": 1, "limits": map[string]any{}} + case "$JS.API.STREAM.LIST": + response = map[string]any{"total": 1, "offset": 0, "limit": 1024, "streams": []any{stream}} + case "$JS.API.STREAM.INFO.ORDERS": + response = stream + case "$JS.API.CONSUMER.LIST.ORDERS": + response = map[string]any{"total": 1, "offset": 0, "limit": 1024, "consumers": []any{consumer}} + case "$JS.API.CONSUMER.INFO.ORDERS.DASHBOARD": + response = consumer + case "$JS.API.STREAM.MSG.GET.ORDERS": + var getRequest struct { + Sequence uint64 `json:"seq"` + } + if err := json.Unmarshal(request, &getRequest); err != nil { + return nil, err + } + response = map[string]any{ + "message": map[string]any{ + "subject": "orders.created", "seq": getRequest.Sequence, "data": []byte("ok"), "time": "2026-08-14T00:00:00Z", + }, + } + default: + return nil, fmt.Errorf("unexpected JetStream API subject %s", subject) + } + return json.Marshal(response) +} + +func writeFakeMessage(writer io.Writer, subject, sid string, payload []byte) error { + _, err := fmt.Fprintf(writer, "MSG %s %s %d\r\n%s\r\n", subject, sid, len(payload), payload) + return err +} + +func readNATSPayload(reader *bufio.Reader, fields []string) ([]byte, error) { + if len(fields) == 0 { + return nil, nil + } + payloadBytes, err := strconv.Atoi(fields[len(fields)-1]) + if err != nil || payloadBytes < 0 { + return nil, errors.New("invalid NATS payload length") + } + data := make([]byte, payloadBytes+2) + if _, err := io.ReadFull(reader, data); err != nil { + return nil, err + } + if !bytes.HasSuffix(data, []byte("\r\n")) { + return nil, errors.New("invalid NATS payload terminator") + } + return data[:payloadBytes], nil +} diff --git a/agents/scripts/driver_release_packages_test.py b/agents/scripts/driver_release_packages_test.py index 79ec082621..6d57f7732e 100644 --- a/agents/scripts/driver_release_packages_test.py +++ b/agents/scripts/driver_release_packages_test.py @@ -26,6 +26,8 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) rabbitmq_source.write_bytes(b"\x7fELFtest-rabbitmq-agent") rocketmq_source = release_dir / "dbx-agent-rocketmq-windows-x64.exe" rocketmq_source.write_bytes(b"MZtest-rocketmq-agent") + nats_source = release_dir / "dbx-agent-nats-linux-x64" + nats_source.write_bytes(b"\x7fELFtest-nats-agent") cassandra_source = release_dir / "dbx-agent-cassandra-linux-x64" cassandra_source.write_bytes(b"\x7fELFtest-cassandra-agent") tdengine_source = release_dir / "dbx-agent-tdengine-windows-aarch64.exe" @@ -44,6 +46,7 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) "rabbitmq": "0.1.0", "rocketmq": "0.1.0", "zookeeper": "0.1.0", + "nats": "0.1.0", "cassandra": "0.1.37", "hive": "0.1.43", "tdengine": "0.1.0", @@ -56,6 +59,7 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) versioned_duckdb = release_dir / "dbx-agent-duckdb-0.1.0-macos-aarch64" versioned_rabbitmq = release_dir / "dbx-agent-rabbitmq-0.1.0-linux-x64" versioned_rocketmq = release_dir / "dbx-agent-rocketmq-0.1.0-windows-x64.exe" + versioned_nats = release_dir / "dbx-agent-nats-0.1.0-linux-x64" versioned_cassandra = release_dir / "dbx-agent-cassandra-0.1.37-linux-x64" versioned_tdengine = release_dir / "dbx-agent-tdengine-0.1.0-windows-aarch64.exe" self.assertEqual( @@ -68,6 +72,7 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) versioned_duckdb, versioned_rabbitmq, versioned_rocketmq, + versioned_nats, versioned_tdengine, ], ) @@ -160,6 +165,19 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) } }, }, + "nats": { + "version": "0.1.0", + "label": "NATS", + "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_nats.name}", + "size": versioned_nats.stat().st_size, + } + }, + }, "tdengine": { "version": "0.1.0", "label": "TDengine", @@ -189,6 +207,7 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) release_dir / "dbx-agent-duckdb-0.1.0-macos-aarch64.tar.zst", release_dir / "dbx-agent-rabbitmq-0.1.0-linux-x64.tar.zst", release_dir / "dbx-agent-rocketmq-0.1.0-windows-x64.tar.zst", + release_dir / "dbx-agent-nats-0.1.0-linux-x64.tar.zst", release_dir / "dbx-agent-tdengine-0.1.0-windows-aarch64.tar.zst", ], ) @@ -200,7 +219,8 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) (outputs[4], "duckdb", versioned_duckdb, "native", "macos-aarch64"), (outputs[5], "rabbitmq", versioned_rabbitmq, "native", "linux-x64"), (outputs[6], "rocketmq", versioned_rocketmq, "native", "windows-x64"), - (outputs[7], "tdengine", versioned_tdengine, "native", "windows-aarch64"), + (outputs[7], "nats", versioned_nats, "native", "linux-x64"), + (outputs[8], "tdengine", versioned_tdengine, "native", "windows-aarch64"), ] for output, driver_name, source, artifact_type, platform in package_cases: tar_bytes = subprocess.run( @@ -231,7 +251,8 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) (final_registry["drivers"]["duckdb"]["native"]["macos-aarch64"], outputs[4]), (final_registry["drivers"]["rabbitmq"]["native"]["linux-x64"], outputs[5]), (final_registry["drivers"]["rocketmq"]["native"]["windows-x64"], outputs[6]), - (final_registry["drivers"]["tdengine"]["native"]["windows-aarch64"], outputs[7]), + (final_registry["drivers"]["nats"]["native"]["linux-x64"], outputs[7]), + (final_registry["drivers"]["tdengine"]["native"]["windows-aarch64"], outputs[8]), ] for artifact, output in release_artifacts: self.assertEqual(artifact["url"], f"https://example.com/{output.name}") @@ -247,6 +268,7 @@ def test_builds_java_and_platform_specific_native_driver_tar_zstd_packages(self) versioned_duckdb, versioned_java, versioned_native, + versioned_nats, versioned_rabbitmq, versioned_rocketmq, versioned_tdengine, diff --git a/agents/scripts/validate_agents.py b/agents/scripts/validate_agents.py index 604e470333..e35dd04ddf 100644 --- a/agents/scripts/validate_agents.py +++ b/agents/scripts/validate_agents.py @@ -11,7 +11,7 @@ KOTLIN_FILE_SUFFIXES = (".kt", ".kts") KOTLIN_SCAN_EXCLUDED_PARTS = {".git", ".gradle", "build"} DEFAULT_AGENT_JRE_KEY = "21" -NON_JDBC_AGENT_MODULES = {"mongodb", "etcd", "zookeeper", "kafka", "rocketmq", "rabbitmq"} +NON_JDBC_AGENT_MODULES = {"mongodb", "etcd", "zookeeper", "kafka", "rocketmq", "rabbitmq", "nats"} NATIVE_ONLY_AGENT_MODULES = { "cassandra": "drivers/cassandra-go", "duckdb": "drivers/duckdb", @@ -26,6 +26,7 @@ "rabbitmq": "drivers/rabbitmq", "rocketmq": "drivers/rocketmq", "zookeeper": "drivers/zookeeper", + "nats": "drivers/nats", } 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 509992b80e..27176671b7 100644 --- a/agents/scripts/validate_agents_test.py +++ b/agents/scripts/validate_agents_test.py @@ -149,11 +149,12 @@ def test_versions_include_native_only_modules(self): "rabbitmq", "rocketmq", "zookeeper", + "nats", "tdengine", ): (root / "drivers" / driver).mkdir(parents=True) (root / "versions.json").write_text( - json.dumps({"h2": "0.1.0", "cassandra": "0.1.0", "hive": "0.1.0", "oracle": "0.1.0", "kingbase": "0.1.0", "iotdb": "0.1.0", "neo4j": "0.1.0", "vastbase": "0.1.0", "xugu": "0.1.0", "rabbitmq": "0.1.0", "rocketmq": "0.1.0", "zookeeper": "0.1.0", "tdengine": "0.1.0"}), + json.dumps({"h2": "0.1.0", "cassandra": "0.1.0", "hive": "0.1.0", "oracle": "0.1.0", "kingbase": "0.1.0", "iotdb": "0.1.0", "neo4j": "0.1.0", "vastbase": "0.1.0", "xugu": "0.1.0", "rabbitmq": "0.1.0", "rocketmq": "0.1.0", "zookeeper": "0.1.0", "nats": "0.1.0", "tdengine": "0.1.0"}), encoding="utf-8", ) diff --git a/agents/scripts/version_agent_artifacts.py b/agents/scripts/version_agent_artifacts.py index 3c5a3a3e6c..7e2c1f4c78 100644 --- a/agents/scripts/version_agent_artifacts.py +++ b/agents/scripts/version_agent_artifacts.py @@ -4,7 +4,7 @@ from pathlib import Path -NATIVE_DRIVERS = ("cassandra", "hive", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "duckdb", "rabbitmq", "rocketmq", "zookeeper", "tdengine") +NATIVE_DRIVERS = ("cassandra", "hive", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "duckdb", "rabbitmq", "rocketmq", "zookeeper", "nats", "tdengine") PLATFORMS = ( "macos-aarch64", "macos-x64", diff --git a/agents/versions.json b/agents/versions.json index dcdf0f7d15..9ff8001da7 100644 --- a/agents/versions.json +++ b/agents/versions.json @@ -43,5 +43,6 @@ "kafka": "0.1.10", "rocketmq": "0.1.7", "rabbitmq": "0.1.6", + "nats": "0.1.0", "sqlserver-legacy": "0.1.16" } diff --git a/apps/desktop/public/icons/database/nats.svg b/apps/desktop/public/icons/database/nats.svg new file mode 100644 index 0000000000..f3608050e7 --- /dev/null +++ b/apps/desktop/public/icons/database/nats.svg @@ -0,0 +1,10 @@ + + NATS + + + + + + + + diff --git a/apps/desktop/src/components/connection/ConnectionDialog.vue b/apps/desktop/src/components/connection/ConnectionDialog.vue index 1ff7ba37a3..a04b06e894 100644 --- a/apps/desktop/src/components/connection/ConnectionDialog.vue +++ b/apps/desktop/src/components/connection/ConnectionDialog.vue @@ -66,6 +66,7 @@ import { assertCompleteDatabaseCategories, databaseSelectionForCategory } from " import { loadConnectionPickerView, saveConnectionPickerView, type DbPickerView } from "@/lib/connection/connectionPickerViewPreference"; import { normalizeRocketmqNamesrvAddr } from "@/lib/connection/rocketmqNamesrv"; import { normalizeRabbitmqAddresses } from "@/lib/connection/rabbitmqAddresses"; +import { NATS_DEFAULT_SERVER_URL, natsConnectionTarget, natsServerUrlIsValid } from "@/lib/connection/natsConnection"; import { detectMqUiAuthKind, isMqAuthKindAllowedForSystem, type MqUiAuthKind } from "@/lib/connection/mqAuth"; import { driverInstallProgressChannel, driverInstallProgressPercent, isDriverInstallProgressForOperation, type DriverInstallProgress } from "@/lib/connection/driverInstallProgressUi"; import { requiresSqlServerLegacyCompatibilityComponent, setSqlServerLegacyCompatibilityConfig, sqlServerUsesLegacyCompatibility, SQLSERVER_LEGACY_COMPATIBILITY_DRIVER_KEY } from "@/lib/connection/sqlServerLegacyCompatibility"; @@ -756,6 +757,7 @@ const mqRocketmqNamesrvAddr = ref("127.0.0.1:9876"); const mqRocketmqClusterName = ref(""); const mqRabbitmqAddresses = ref("127.0.0.1:5672"); const mqRabbitmqVirtualHost = ref("/"); +const mqNatsServerUrl = ref(NATS_DEFAULT_SERVER_URL); const mqKafkaBootstrapServers = ref("127.0.0.1:9092"); const mqKafkaZooKeeperServers = ref(""); const mqKafkaSecurityProtocol = ref(MQ_KAFKA_SECURITY_PROTOCOL_AUTO); @@ -784,12 +786,14 @@ const MQ_DRIVER_LABELS: Record = { kafka: "Apache Kafka", rocketmq: "Apache RocketMQ", rabbitmq: "RabbitMQ", + nats: "NATS", }; function mqSystemKindFromProfile(profile: string): MqSystemKind { if (profile === "kafka") return "kafka"; if (profile === "rocketmq") return "rocketmq"; if (profile === "rabbitmq") return "rabbitmq"; + if (profile === "nats") return "nats"; return "pulsar"; } @@ -799,7 +803,7 @@ function syncMqSystemKindFromSelectedType() { } function resolveMqSystemKind(config?: Partial): MqSystemKind { - if (config?.systemKind === "kafka" || config?.systemKind === "rocketmq" || config?.systemKind === "rabbitmq" || config?.systemKind === "pulsar") { + if (config?.systemKind === "kafka" || config?.systemKind === "rocketmq" || config?.systemKind === "rabbitmq" || config?.systemKind === "pulsar" || config?.systemKind === "nats") { return config.systemKind; } return mqSystemKindFromProfile(selectedType.value); @@ -1192,6 +1196,7 @@ const driverProfiles: Record< kafka: { type: "mq", port: 9092, user: "", label: "Apache Kafka", icon: "kafka", host: "127.0.0.1" }, rocketmq: { type: "mq", port: 9876, user: "", label: "Apache RocketMQ", icon: "rocketmq", host: "127.0.0.1" }, rabbitmq: { type: "mq", port: 5672, user: "", label: "RabbitMQ", icon: "rabbitmq", host: "127.0.0.1" }, + nats: { type: "mq", port: 4222, user: "", label: "NATS", icon: "nats", host: "127.0.0.1" }, nacos: { type: "nacos", port: 8848, user: "nacos", label: "Nacos", icon: "nacos", host: "127.0.0.1" }, consul: { type: "consul", port: 8500, user: "", label: "Consul", icon: "consul", host: "127.0.0.1" }, mqtt: { type: "mqtt", port: 1883, user: "", label: "MQTT", icon: "mqtt", host: "127.0.0.1" }, @@ -1229,6 +1234,7 @@ function profileForConfig(config: ConnectionConfig) { if (kind === "kafka") return "kafka"; if (kind === "rocketmq") return "rocketmq"; if (kind === "rabbitmq") return "rabbitmq"; + if (kind === "nats") return "nats"; return "mq"; } if (config.db_type === "dameng") return "dm"; @@ -1278,7 +1284,7 @@ function resetMqFields(config?: Partial) { const jaasConfig = mqExtraPropertyString(extra, "sasl.jaas.config"); mqSystemKind.value = systemKind; const storedAdminUrl = config?.adminUrl?.trim() || (config ? mqExtraString(config as Record, "admin_url").trim() : ""); - mqAdminUrl.value = storedAdminUrl || (systemKind === "kafka" || systemKind === "rocketmq" || systemKind === "rabbitmq" ? "" : "http://127.0.0.1:8080"); + mqAdminUrl.value = storedAdminUrl || (systemKind === "kafka" || systemKind === "rocketmq" || systemKind === "rabbitmq" || systemKind === "nats" ? "" : "http://127.0.0.1:8080"); mqKafkaConnectionSource.value = resolveMqKafkaConnectionSource(extra); mqKafkaBootstrapServers.value = mqExtraString(extra, "bootstrapServers") || "127.0.0.1:9092"; mqKafkaZooKeeperServers.value = mqExtraString(extra, "zookeeperServers"); @@ -1286,6 +1292,7 @@ function resetMqFields(config?: Partial) { mqRocketmqClusterName.value = mqExtraString(extra, "clusterName") || mqExtraString(extra, "cluster_name"); mqRabbitmqAddresses.value = mqExtraString(extra, "addresses") || "127.0.0.1:5672"; mqRabbitmqVirtualHost.value = mqExtraString(extra, "virtualHost") || "/"; + mqNatsServerUrl.value = config?.serverUrl?.trim() || NATS_DEFAULT_SERVER_URL; mqKafkaSecurityProtocol.value = mqExtraString(extra, "securityProtocol") || MQ_KAFKA_SECURITY_PROTOCOL_AUTO; mqKafkaSaslMechanism.value = mqExtraString(extra, "saslMechanism") || "PLAIN"; mqKafkaKerberosPrincipal.value = parseJaasStringProperty(jaasConfig, "principal"); @@ -1295,15 +1302,25 @@ function resetMqFields(config?: Partial) { mqTlsSkipVerify.value = !!config?.tlsSkipVerify; mqPinnedVersion.value = pinnedVersionToSelection(config?.pinnedVersion); const auth = (config?.auth || { kind: "none" }) as MqAuth; - mqAuthKind.value = detectMqUiAuthKind({ - systemKind, - authKind: auth.kind, - saslMechanism: mqKafkaSaslMechanism.value, - jaasConfig, - }); - mqToken.value = auth.token || ""; - mqBasicUsername.value = auth.username || ""; - mqBasicPassword.value = auth.password || ""; + if (systemKind === "nats") { + const token = config?.token || auth.token || ""; + const username = config?.username || auth.username || ""; + const password = config?.password || auth.password || ""; + mqAuthKind.value = token ? "token" : username || auth.kind === "basic" ? "basic" : "none"; + mqToken.value = token; + mqBasicUsername.value = username; + mqBasicPassword.value = password; + } else { + mqAuthKind.value = detectMqUiAuthKind({ + systemKind, + authKind: auth.kind, + saslMechanism: mqKafkaSaslMechanism.value, + jaasConfig, + }); + mqToken.value = auth.token || ""; + mqBasicUsername.value = auth.username || ""; + mqBasicPassword.value = auth.password || ""; + } mqApiKeyHeader.value = auth.header || "Authorization"; mqApiKeyValue.value = auth.value || ""; mqOauthIssuerUrl.value = auth.issuerUrl || ""; @@ -1341,6 +1358,14 @@ function defaultMqFieldsForProfile(profile: string): Partial | un extra: { addresses: "127.0.0.1:5672", virtualHost: "/" }, }; } + if (profile === "nats") { + return { + systemKind: "nats", + adminUrl: "", + auth: { kind: "none" }, + serverUrl: NATS_DEFAULT_SERVER_URL, + }; + } return undefined; } @@ -1373,6 +1398,11 @@ watch(mqSystemKind, (kind) => { if (!isMqAuthKindAllowedForSystem(kind, mqAuthKind.value)) mqAuthKind.value = "none"; return; } + if (kind === "nats") { + if (!mqNatsServerUrl.value.trim()) mqNatsServerUrl.value = NATS_DEFAULT_SERVER_URL; + if (!isMqAuthKindAllowedForSystem(kind, mqAuthKind.value)) mqAuthKind.value = "none"; + return; + } if (!mqAdminUrl.value.trim()) mqAdminUrl.value = "http://127.0.0.1:8080"; }); @@ -1685,6 +1715,40 @@ function buildMqAdminConfig(): MqAdminConfig { }; } + if (systemKind === "nats") { + const target = natsConnectionTarget(mqNatsServerUrl.value); + if (mqAuthKind.value === "basic") { + return { + systemKind: "nats", + adminUrl: "", + // `auth.kind` is non-secret metadata used by the secret store to + // restore the canonical top-level NATS credentials after restart. + auth: { kind: "basic" }, + serverUrl: target.serverUrl, + username: requireMqField(mqBasicUsername.value, "NATS password auth requires a username"), + password: requireMqField(mqBasicPassword.value, "NATS password auth requires a password"), + tlsSkipVerify: mqTlsSkipVerify.value || undefined, + }; + } + if (mqAuthKind.value === "token") { + return { + systemKind: "nats", + adminUrl: "", + auth: { kind: "token" }, + serverUrl: target.serverUrl, + token: requireMqField(mqToken.value, "NATS token auth requires a token"), + tlsSkipVerify: mqTlsSkipVerify.value || undefined, + }; + } + return { + systemKind: "nats", + adminUrl: "", + auth: { kind: "none" }, + serverUrl: target.serverUrl, + tlsSkipVerify: mqTlsSkipVerify.value || undefined, + }; + } + if (systemKind === "rabbitmq") { const addresses = normalizeRabbitmqAddresses(mqRabbitmqAddresses.value); const extra: Record = { @@ -2900,6 +2964,7 @@ const dbOptions: DbOption[] = [ { value: "kafka", label: "Apache Kafka" }, { value: "rocketmq", label: "Apache RocketMQ" }, { value: "rabbitmq", label: "RabbitMQ" }, + { value: "nats", label: "NATS" }, { value: "mqtt", label: "MQTT" }, { value: "nacos", label: "Nacos" }, { value: "consul", label: "Consul" }, @@ -2958,7 +3023,7 @@ const dbCategoryDefinitions: Array<{ { key: "mq", titleKey: "connection.databaseCategoryMq", - optionValues: ["mq", "kafka", "rocketmq", "rabbitmq", "mqtt"], + optionValues: ["mq", "kafka", "rocketmq", "rabbitmq", "nats", "mqtt"], }, { key: "registry_config", @@ -3440,6 +3505,7 @@ const hasRequiredConnectionTarget = computed(() => { if (mqSystemKind.value === "kafka") return mqKafkaConnectionSource.value === "zookeeper" ? !!mqKafkaZooKeeperServers.value.trim() : !!mqKafkaBootstrapServers.value.trim(); if (mqSystemKind.value === "rocketmq") return !!mqRocketmqNamesrvAddr.value.trim(); if (mqSystemKind.value === "rabbitmq") return !!mqRabbitmqAddresses.value.trim(); + if (mqSystemKind.value === "nats") return natsServerUrlIsValid(mqNatsServerUrl.value); return !!mqAdminUrl.value.trim(); } if (form.value.db_type === "zookeeper") return !!(form.value.host || form.value.connection_string || connectionUrlInput.value.trim()); @@ -3831,6 +3897,11 @@ function connectionConfigForSubmit(id: string, generatedName = ""): ConnectionCo } else if (mqConfig.systemKind === "rabbitmq") { const extra = mqExtraRecord(mqConfig); applyMqRabbitmqAddresses(config, mqExtraString(extra, "addresses")); + } else if (mqConfig.systemKind === "nats") { + const target = natsConnectionTarget(mqConfig.serverUrl || NATS_DEFAULT_SERVER_URL); + config.host = target.host; + config.port = target.port; + config.ssl = target.tls; } else { applyMqAdminUrl(config, mqConfig.adminUrl); } @@ -6083,6 +6154,15 @@ function openExternalUrl(url: string) { +