diff --git a/.github/scripts/bump-agent-versions.mjs b/.github/scripts/bump-agent-versions.mjs index bc2705e8e..94993a949 100644 --- a/.github/scripts/bump-agent-versions.mjs +++ b/.github/scripts/bump-agent-versions.mjs @@ -48,9 +48,10 @@ const nativeDriverDirectories = { duckdb: "duckdb", oracle: "oracle-go", kingbase: "kingbase-go", + vastbase: "vastbase-go", rabbitmq: "rabbitmq", }; -const nativeDriverModules = new Set(["duckdb", "oracle", "xugu", "kingbase", "rabbitmq"]); +const nativeDriverModules = new Set(["duckdb", "oracle", "xugu", "kingbase", "vastbase", "rabbitmq"]); function resolveAgentModule(moduleName, { legacyStandaloneModules, moduleExists, readModuleFile }) { let checkDir = null; diff --git a/.github/scripts/bump-agent-versions.test.mjs b/.github/scripts/bump-agent-versions.test.mjs index e2c1b5e30..1f61f0c3d 100644 --- a/.github/scripts/bump-agent-versions.test.mjs +++ b/.github/scripts/bump-agent-versions.test.mjs @@ -45,6 +45,18 @@ test("bumps the native RabbitMQ agent from its Go directory", () => { assert.equal(result.versions.rabbitmq, "0.1.1"); }); +test("bumps the native Vastbase agent from its independent Go directory", () => { + const result = evaluateAgentVersionBump({ + versions: { vastbase: "0.1.37" }, + changedFiles: ["agents/drivers/vastbase-go/main.go"], + moduleExists: (path) => path === "agents/drivers/vastbase-go", + readModuleFile: () => "", + }); + + assert.equal(result.versions.vastbase, "0.1.38"); + assert.deepEqual(result.nativeModules, ["vastbase"]); +}); + test("builds a manually versioned module even without runtime file changes", () => { const result = evaluateAgentVersionBump({ versions: { duckdb: "0.1.1" }, diff --git a/.github/scripts/label-pull-request.mjs b/.github/scripts/label-pull-request.mjs index 67281dcdd..f903a7ead 100644 --- a/.github/scripts/label-pull-request.mjs +++ b/.github/scripts/label-pull-request.mjs @@ -58,6 +58,7 @@ const DRIVER_DATABASE_ALIASES = { rabbitmq: "mq", rocketmq: "mq", "sqlserver-legacy": "sqlserver", + "vastbase-go": "vastbase", }; const DIALECT_DATABASE_ALIASES = { diff --git a/.github/scripts/label-pull-request.test.mjs b/.github/scripts/label-pull-request.test.mjs index 5b30bc453..ed89fb522 100644 --- a/.github/scripts/label-pull-request.test.mjs +++ b/.github/scripts/label-pull-request.test.mjs @@ -25,6 +25,7 @@ const knownDatabaseTypes = new Set([ "redis", "sqlite", "sqlserver", + "vastbase", ]); test("labels a desktop MySQL UI fix", () => { @@ -54,11 +55,12 @@ test("maps agent and dialect paths to existing database types", () => { assert.deepEqual( inferDatabaseTypes([ "agents/drivers/oracle-go/go.mod", + "agents/drivers/vastbase-go/go.mod", "agents/drivers/kafka/build.gradle", "plugins/dialects/postgresql.yaml", "plugins/dialects/oceanbase.yaml", ], knownDatabaseTypes), - ["mq", "oceanbase-oracle", "oracle", "postgres"], + ["mq", "oceanbase-oracle", "oracle", "postgres", "vastbase"], ); }); diff --git a/.github/scripts/reuse-agent-release-assets.mjs b/.github/scripts/reuse-agent-release-assets.mjs index 74426848c..8a342adf9 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", "rabbitmq"]); +const NATIVE_MODULES = new Set(["duckdb", "oracle", "xugu", "kingbase", "vastbase", "rabbitmq"]); 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 e67ff2613..b99b7655b 100644 --- a/.github/scripts/reuse-agent-release-assets.test.mjs +++ b/.github/scripts/reuse-agent-release-assets.test.mjs @@ -44,16 +44,16 @@ test("collects complete reusable Java, native, and JRE assets", () => { test("rejects an incomplete reusable native platform set", () => { const native = Object.fromEntries( - platforms.slice(1).map((platform, index) => [platform, artifact(`dbx-agent-duckdb-0.1.2-${platform}.tar.zst`, String(index + 1))]), + platforms.slice(1).map((platform, index) => [platform, artifact(`dbx-agent-vastbase-0.1.38-${platform}.tar.zst`, String(index + 1))]), ); - const registry = { drivers: { duckdb: { version: "0.1.2", native } }, jres: {} }; + const registry = { drivers: { vastbase: { version: "0.1.38", native } }, jres: {} }; assert.throws( () => collectReusableAssetPlan({ registry, release: releaseFor(Object.values(native)), - versions: { duckdb: "0.1.2" }, - modules: ["duckdb"], + versions: { vastbase: "0.1.38" }, + modules: ["vastbase"], reuseJre: false, }), /missing=macos-aarch64/, diff --git a/.github/workflows/agents-release.yml b/.github/workflows/agents-release.yml index 35afd884b..3e987555a 100644 --- a/.github/workflows/agents-release.yml +++ b/.github/workflows/agents-release.yml @@ -316,6 +316,46 @@ jobs: name: kingbase-native path: "release-native/dbx-agent-kingbase-*" + build-vastbase-native: + needs: [bump-versions] + if: ${{ contains(fromJSON(needs.bump-versions.outputs.native_modules), 'vastbase') }} + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.22.x" + - name: Test Vastbase native agent + working-directory: agents/drivers/vastbase-go + run: go test ./... + - name: Cross-compile Vastbase native agent + shell: bash + run: | + mkdir -p release-native + cd agents/drivers/vastbase-go + 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-vastbase-${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: vastbase-native + path: "release-native/dbx-agent-vastbase-*" + build-duckdb-native: needs: [bump-versions] if: ${{ contains(fromJSON(needs.bump-versions.outputs.native_modules), 'duckdb') }} @@ -531,7 +571,7 @@ jobs: retention-days: 1 release: - needs: [bump-versions, commit-versions, build-agents, build-oracle-native, build-xugu-native, build-rabbitmq-native, build-kingbase-native, build-duckdb-native, build-jre, reuse-previous-assets] + needs: [bump-versions, commit-versions, build-agents, build-oracle-native, build-xugu-native, build-rabbitmq-native, build-kingbase-native, build-vastbase-native, build-duckdb-native, build-jre, reuse-previous-assets] if: ${{ always() && !contains(needs.*.result, 'failure') && !contains(needs.*.result, 'cancelled') }} runs-on: ubuntu-latest steps: @@ -649,6 +689,7 @@ jobs: local name="$1" case "$name" in kingbase) echo "人大金仓 KingbaseES" ;; + vastbase) echo "Vastbase" ;; duckdb) echo "DuckDB" ;; xugu) echo "虚谷 XuguDB" ;; rabbitmq) echo "RabbitMQ" ;; @@ -715,7 +756,7 @@ jobs: [ -n "$DRIVERS" ] && DRIVERS="${DRIVERS},"$'\n' DRIVERS="${DRIVERS}$(generate_jar_entry "$name" "$label" "$f" "$jre_key" "$version" "$external_driver" "$native_json")" done - for name in oracle xugu kingbase duckdb rabbitmq; do + for name in oracle xugu kingbase vastbase duckdb rabbitmq; do version=$(get_module_version "$name") [ -f "release/dbx-agent-${name}-${version}.jar" ] && continue native_json=$(generate_native_platforms "$name" "$version") @@ -780,6 +821,7 @@ jobs: local name="$1" case "$name" in kingbase) echo "人大金仓 KingbaseES" ;; + vastbase) echo "Vastbase" ;; duckdb) echo "DuckDB" ;; oracle) echo "Oracle" ;; xugu) echo "虚谷 XuguDB" ;; @@ -806,6 +848,8 @@ jobs: LOG_PATH="agents/drivers/oracle-go/" elif [ "$name" = "kingbase" ]; then LOG_PATH="agents/drivers/kingbase-go/" + elif [ "$name" = "vastbase" ]; then + LOG_PATH="agents/drivers/vastbase-go/" elif [ -d "agents/drivers/$name" ]; then LOG_PATH="agents/drivers/$name/" else diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 22b760e9c..fd06d1b01 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -538,6 +538,10 @@ jobs: run: go test ./... working-directory: agents/drivers/rabbitmq + - name: Vastbase native agent tests + run: go test ./... + working-directory: agents/drivers/vastbase-go + - name: Oracle native agent build run: CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -trimpath -ldflags="-s -w" -o /tmp/dbx-agent-oracle-linux-x64 . working-directory: agents/drivers/oracle-go @@ -550,6 +554,10 @@ jobs: run: CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -trimpath -ldflags="-s -w" -o /tmp/dbx-agent-rabbitmq-linux-x64 . working-directory: agents/drivers/rabbitmq + - name: Vastbase native agent build + run: CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -trimpath -ldflags="-s -w" -o /tmp/dbx-agent-vastbase-linux-x64 . + working-directory: agents/drivers/vastbase-go + - name: RabbitMQ native agent integration tests shell: bash working-directory: agents/drivers/rabbitmq diff --git a/agents/README.md b/agents/README.md index 6f24213c3..f4fe28175 100644 --- a/agents/README.md +++ b/agents/README.md @@ -13,7 +13,7 @@ Each agent runs as a standalone process and communicates with DBX via stdin/stdo | access | Microsoft Access | UCanAccess | | dameng | 达梦 DM8 | DM JDBC | | kingbase | 人大金仓 KingbaseES | gokb Go native agent | -| vastbase | Vastbase | Vastbase JDBC | +| vastbase | Vastbase | openGauss Go native agent | | uxdb | UXDB | UXDB JDBC | | goldendb | GoldenDB | MySQL Connector/J | | databend | Databend | Databend JDBC | @@ -75,7 +75,7 @@ Set `DBX_AGENT_JDBC_POOL_ENABLED=false` for a runtime-level compatibility fallba For new agents, prefer a **native (Go or Rust) driver** over a Java/JDBC agent whenever a mature, license-compatible native driver is available. Native agents ship as a single self-contained executable with no JRE, which significantly reduces memory footprint and startup time — the JVM baseline that every Java agent pays even when idle is avoided entirely. -- **Native (C++/Go/Rust)** — preferred when a usable native driver exists. See `drivers/duckdb`, `drivers/oracle-go` (go-ora), `drivers/kingbase-go` (gokb), `drivers/xugu`, and `drivers/rabbitmq` (amqp091-go) as reference implementations. No JRE download or management is needed. +- **Native (C++/Go/Rust)** — preferred when a usable native driver exists. See `drivers/duckdb`, `drivers/oracle-go` (go-ora), `drivers/kingbase-go` (gokb), `drivers/vastbase-go` (openGauss connector), `drivers/xugu`, and `drivers/rabbitmq` (amqp091-go) as reference implementations. No JRE download or management is needed. - **Java/JDBC** — the default fallback when only a JDBC driver exists for the database, or when the native driver is immature or unmaintained. Most agents still fall in this category. Native agents implement the same JSON-RPC contract and `versions.json` registration as Java agents; they ship an `agent` executable instead of `agent.jar`. If both native and Java source implementations exist for the same database, publish only the native artifact unless the Java variant has a separately registered compatibility profile, such as `oracle-legacy` / `oracle-10g`. @@ -88,11 +88,12 @@ Requires JDK 21 (Gradle toolchain auto-downloads if needed). ./gradlew shadowJar (cd drivers/oracle-go && go build -o agent .) (cd drivers/kingbase-go && go build -o agent .) +(cd drivers/vastbase-go && go build -o agent .) (cd drivers/xugu && go build -o agent .) (cd drivers/rabbitmq && go build -o agent .) ``` -Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/oracle-go`, `drivers/kingbase-go`, `drivers/xugu`, and `drivers/rabbitmq`. +Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/oracle-go`, `drivers/kingbase-go`, `drivers/vastbase-go`, `drivers/xugu`, and `drivers/rabbitmq`. ### Local DBX Runtime Test diff --git a/agents/README.zh-CN.md b/agents/README.zh-CN.md index 7e5b8d5a4..fff0b60b8 100644 --- a/agents/README.zh-CN.md +++ b/agents/README.zh-CN.md @@ -13,7 +13,7 @@ DBX 的 Agent 驱动 —— 通过 JDBC 和原生数据库驱动支持各种数 | access | Microsoft Access | UCanAccess | | dameng | 达梦 DM8 | DM JDBC | | kingbase | 人大金仓 KingbaseES | gokb Go 原生 agent | -| vastbase | Vastbase | Vastbase JDBC | +| vastbase | Vastbase | openGauss Go 原生 agent | | uxdb | 优炫 UXDB | UXDB JDBC | | goldendb | GoldenDB | MySQL Connector/J | | databend | Databend | Databend JDBC | @@ -75,7 +75,7 @@ HikariCP 会直接打进启用连接池的 Agent JAR。已经使用 DBX 托管 J 对于新 agent,只要存在成熟、许可证兼容的原生驱动,优先选择**原生(Go 或 Rust)驱动**而非 Java/JDBC agent。原生 agent 以单一自包含可执行文件发布,无需 JRE,可显著降低内存占用和启动时间 —— 完全避开 Java agent 即便空闲也要付出的 JVM 基线开销。 -- **原生(Go/Rust)** —— 存在可用原生驱动时首选。参考 `drivers/oracle-go`(go-ora)、`drivers/kingbase-go`(gokb)、`drivers/xugu` 和 `drivers/rabbitmq`(amqp091-go)。无需 JRE 下载与管理。 +- **原生(Go/Rust)** —— 存在可用原生驱动时首选。参考 `drivers/oracle-go`(go-ora)、`drivers/kingbase-go`(gokb)、`drivers/vastbase-go`(openGauss connector)、`drivers/xugu` 和 `drivers/rabbitmq`(amqp091-go)。无需 JRE 下载与管理。 - **Java/JDBC** —— 当某数据库只有 JDBC 驱动,或原生驱动不成熟、缺乏维护时的默认兜底方案。多数 agent 仍属此类。 原生 agent 实现与 Java agent 相同的 JSON-RPC 契约和 `versions.json` 登记;它发布的是 `agent` 可执行文件而非 `agent.jar`。若同一数据库同时保留原生和 Java 源码实现,默认只发布原生产物;只有 Java 变体以独立兼容配置登记时才同时发布,例如 `oracle-legacy` / `oracle-10g`。 @@ -88,11 +88,12 @@ HikariCP 会直接打进启用连接池的 Agent JAR。已经使用 DBX 托管 J ./gradlew shadowJar (cd drivers/oracle-go && go build -o agent .) (cd drivers/kingbase-go && go build -o agent .) +(cd drivers/vastbase-go && go build -o agent .) (cd drivers/xugu && go build -o agent .) (cd drivers/rabbitmq && go build -o agent .) ``` -产物 JAR 在 `drivers/{module}/build/libs/`。原生 agent 从 `drivers/oracle-go`、`drivers/kingbase-go`、`drivers/xugu` 和 `drivers/rabbitmq` 构建。 +产物 JAR 在 `drivers/{module}/build/libs/`。原生 agent 从 `drivers/oracle-go`、`drivers/kingbase-go`、`drivers/vastbase-go`、`drivers/xugu` 和 `drivers/rabbitmq` 构建。 ### 本地 DBX 运行时测试 diff --git a/agents/build.gradle b/agents/build.gradle index 181160ece..fe86b937d 100644 --- a/agents/build.gradle +++ b/agents/build.gradle @@ -9,7 +9,7 @@ def pooledJdbcProjects = [ 'firebird', 'gbase8a', 'gbase8s', 'goldendb', 'h2', 'h2-legacy', 'highgo', 'hive', 'informix', 'iotdb', 'iris', 'kylin', 'neo4j', 'oceanbase-oracle', 'oscar', 'saphana', 'snowflake', 'spark', 'sqlserver-legacy', 'sundb', 'tdengine', 'teradata', 'trino', 'uxdb', - 'vastbase', 'vertica', 'yashandb' + 'vertica', 'yashandb' ] as Set def agentProjects = subprojects.findAll { !infrastructureProjects.contains(it.name) } def jdbcAgentProjects = agentProjects.findAll { !legacyStandaloneProjects.contains(it.name) } diff --git a/agents/drivers/vastbase-go/README.md b/agents/drivers/vastbase-go/README.md new file mode 100644 index 000000000..19bebbd3c --- /dev/null +++ b/agents/drivers/vastbase-go/README.md @@ -0,0 +1,36 @@ +# Vastbase Native Agent + +This module implements the DBX agent protocol for Vastbase with the pure-Go +`openGauss-connector-go-pq` driver. + +## Build + +```bash +go test ./... +CGO_ENABLED=0 go build -trimpath -ldflags="-s -w" -o agent . +``` + +## Local DBX Test + +Build the binary, then copy it into DBX's installed Vastbase driver directory: + +```bash +mkdir -p ~/.dbx/agents/drivers/vastbase +cp agent ~/.dbx/agents/drivers/vastbase/agent +chmod +x ~/.dbx/agents/drivers/vastbase/agent +``` + +DBX prefers `agent` over `agent.jar`. Remove the native binary to restore a +previously installed JDBC agent. + +## Integration Test + +Set `VASTBASE_TEST_HOST`, `VASTBASE_TEST_PORT`, `VASTBASE_TEST_DATABASE`, +`VASTBASE_TEST_USERNAME`, and `VASTBASE_TEST_PASSWORD`, then run: + +```bash +go test -run '^TestVastbaseIntegration$' -count=1 ./... +``` + +The benchmark harness under `bench/` compares this native agent with the +Vastbase JDBC 2.11v and 2.15v agents. diff --git a/agents/drivers/vastbase-go/bench/README.md b/agents/drivers/vastbase-go/bench/README.md new file mode 100644 index 000000000..8a76f9166 --- /dev/null +++ b/agents/drivers/vastbase-go/bench/README.md @@ -0,0 +1,106 @@ +# Vastbase Agent benchmark + +This benchmark compares the same DBX JSON-RPC workload through: + +- Vastbase JDBC `2.11v` (current DBX baseline) +- Vastbase JDBC `2.15v` +- openGauss Go connector `v1.0.8` + +It keeps one physical database connection per Agent session and reports startup, +connection/authentication, steady-state latency, throughput, RSS, and artifact size. + +## Build candidates + +```bash +mkdir -p /tmp/dbx-vastbase-bench + +go build -o /tmp/dbx-vastbase-bench/vastbase-go ./drivers/vastbase-go +go build -o /tmp/dbx-vastbase-bench/agent-compare ./drivers/vastbase-go/bench +``` + +Run these commands from `agents/`. Supply the archived JDBC `2.11v` and `2.15v` +agent JARs through `JDBC_211_AGENT_JAR` and `JDBC_215_AGENT_JAR`; the production +Vastbase module no longer builds or ships the JDBC agent. + +## Direct connector probe + +Use the direct probe to separate openGauss connector row-decoding cost from the +Agent JSON-RPC path. It pins one physical connection per worker and runs the same +1,000-row decode query used by `decode_rows`. + +```bash +go build -o /tmp/dbx-vastbase-bench/vastbase-direct ./drivers/vastbase-go/bench/direct + +DBX_TEST_PASSWORD='secret' \ +VASTBASE_HOST=127.0.0.1 \ +VASTBASE_PORT=5432 \ +VASTBASE_DATABASE=postgres \ +VASTBASE_USERNAME=vastbase \ +BENCH_MODE=collect \ +BENCH_CONCURRENCY=32 \ +BENCH_SECONDS=4 \ +/tmp/dbx-vastbase-bench/vastbase-direct +``` + +Set `BENCH_MODE=marshal` to include `encoding/json` serialization of the collected +result. The probe is diagnostic evidence for the Go connector path; it is not a +replacement for the JDBC-vs-Go Agent benchmark. + +## Startup-only benchmark + +Startup does not require a database server: + +```bash +JDBC_211_AGENT_JAR=/tmp/dbx-vastbase-bench/vastbase-jdbc-2.11v.jar \ +JDBC_215_AGENT_JAR=/tmp/dbx-vastbase-bench/vastbase-jdbc-2.15v.jar \ +GO_AGENT=/tmp/dbx-vastbase-bench/vastbase-go \ +BENCH_PHASES=startup \ +/tmp/dbx-vastbase-bench/agent-compare > /tmp/dbx-vastbase-bench/startup.ndjson +``` + +## Live G100/V100 benchmark + +```bash +JDBC_211_AGENT_JAR=/tmp/dbx-vastbase-bench/vastbase-jdbc-2.11v.jar \ +JDBC_215_AGENT_JAR=/tmp/dbx-vastbase-bench/vastbase-jdbc-2.15v.jar \ +GO_AGENT=/tmp/dbx-vastbase-bench/vastbase-go \ +VASTBASE_HOST=127.0.0.1 \ +VASTBASE_PORT=5432 \ +VASTBASE_DATABASE=postgres \ +VASTBASE_USERNAME=vastbase \ +VASTBASE_PASSWORD='secret' \ +VASTBASE_SERVER='G100-V3.0.9-test' \ +/tmp/dbx-vastbase-bench/agent-compare > /tmp/dbx-vastbase-bench/live.ndjson +``` + +Defaults: + +- phases: `startup,connect,query` +- workloads: `select_literal,decode_rows,page_rows,list_tables` +- rounds: `3` +- measured duration: `4s` per workload/agent/concurrency/round +- concurrency: `1,8,32` +- decoded/page rows: `1000` + +The following variables can override the defaults: + +- `BENCH_PHASES` +- `BENCH_WORKLOADS` +- `BENCH_ROUNDS` +- `BENCH_SECONDS` +- `BENCH_CONCURRENCIES` +- `BENCH_LITERAL_SQL` +- `BENCH_DECODE_SQL` +- `BENCH_DECODE_ROWS` +- `BENCH_PAGE_SQL` +- `BENCH_PAGE_ROWS` +- `BENCH_SCHEMA` +- `VASTBASE_SSL` +- `VASTBASE_URL_PARAMS` +- `VASTBASE_CONNECTION_STRING` +- `VASTBASE_CA_CERT_PATH` +- `VASTBASE_CLIENT_CERT_PATH` +- `VASTBASE_CLIENT_KEY_PATH` + +Do not use a PostgreSQL/openGauss mock to make a Vastbase performance claim. +Record the exact G100/V100 edition and server version from the emitted metadata. diff --git a/agents/drivers/vastbase-go/bench/agent_compare.go b/agents/drivers/vastbase-go/bench/agent_compare.go new file mode 100644 index 000000000..cc555b896 --- /dev/null +++ b/agents/drivers/vastbase-go/bench/agent_compare.go @@ -0,0 +1,867 @@ +package main + +import ( + "bufio" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "os/exec" + "runtime" + "sort" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" +) + +type agentSpec struct { + Name string + Command []string + ArtifactPath string +} + +type agentProcess struct { + command *exec.Cmd + stdin io.WriteCloser + pending sync.Map + writeMu sync.Mutex + nextID atomic.Int64 + reader *bufio.Scanner +} + +type agentResponse struct { + ID int64 `json:"id"` + Result json.RawMessage `json:"result"` + Error *struct { + Message string `json:"message"` + } `json:"error"` +} + +type benchmarkMetadata struct { + Type string `json:"type"` + Server string `json:"server,omitempty"` + DatabaseVersions map[string]string `json:"database_versions,omitempty"` + GOOS string `json:"goos"` + GOARCH string `json:"goarch"` + Phases []string `json:"phases"` + Rounds int `json:"rounds"` + DurationSeconds int `json:"duration_seconds"` + Concurrencies []int `json:"concurrencies"` + StartupWarmups int `json:"startup_warmups"` + StartupIterations int `json:"startup_iterations"` + ConnectWarmups int `json:"connect_warmups"` + ConnectIterations int `json:"connect_iterations"` + QueryWarmups int `json:"query_warmups"` + Workloads []string `json:"workloads"` +} + +type benchmarkResult struct { + Type string `json:"type"` + Server string `json:"server,omitempty"` + Agent string `json:"agent"` + Workload string `json:"workload"` + Round int `json:"round"` + Concurrency int `json:"concurrency"` + Operations int64 `json:"operations"` + Errors int64 `json:"errors"` + DurationMS float64 `json:"duration_ms"` + QPS float64 `json:"qps"` + MeanMS float64 `json:"mean_ms"` + P50MS float64 `json:"p50_ms"` + P95MS float64 `json:"p95_ms"` + P99MS float64 `json:"p99_ms"` + ReadyRSSKB int64 `json:"ready_rss_kb,omitempty"` + OneSessionKB int64 `json:"one_session_rss_kb,omitempty"` + AllSessionsKB int64 `json:"all_sessions_rss_kb,omitempty"` + PeakRSSKB int64 `json:"peak_rss_kb,omitempty"` + ArtifactBytes int64 `json:"artifact_bytes,omitempty"` +} + +type runningAgent struct { + spec agentSpec + process *agentProcess + readyRSSKB int64 + oneSessionRSSKB int64 + allSessionsRSS int64 +} + +type workload struct { + Name string + Method string + Parameters func(worker int) map[string]any + Cleanup func(*agentProcess, json.RawMessage, int) error +} + +func main() { + agents := []agentSpec{ + { + Name: "jdbc-2.11v", + Command: jdbcAgentCommand(requiredEnv("JDBC_211_AGENT_JAR")), + ArtifactPath: requiredEnv("JDBC_211_AGENT_JAR"), + }, + { + Name: "jdbc-2.15v", + Command: jdbcAgentCommand(requiredEnv("JDBC_215_AGENT_JAR")), + ArtifactPath: requiredEnv("JDBC_215_AGENT_JAR"), + }, + { + Name: "go-v1.0.8", + Command: []string{requiredEnv("GO_AGENT")}, + ArtifactPath: requiredEnv("GO_AGENT"), + }, + } + for _, agent := range agents { + if _, err := os.Stat(agent.ArtifactPath); err != nil { + panic(fmt.Errorf("stat %s artifact %s: %w", agent.Name, agent.ArtifactPath, err)) + } + } + + phases := envStrings("BENCH_PHASES", []string{"startup", "connect", "query"}) + rounds := envInt("BENCH_ROUNDS", 3) + durationSeconds := envInt("BENCH_SECONDS", 4) + concurrencies := envInts("BENCH_CONCURRENCIES", []int{1, 8, 32}) + startupWarmups := envNonNegativeInt("BENCH_STARTUP_WARMUPS", 2) + startupIterations := envInt("BENCH_STARTUPS", 20) + connectWarmups := envNonNegativeInt("BENCH_CONNECT_WARMUPS", 3) + connectIterations := envInt("BENCH_CONNECTS", 30) + queryWarmups := envNonNegativeInt("BENCH_QUERY_WARMUPS", 20) + workloadNames := envStrings("BENCH_WORKLOADS", []string{"select_literal", "decode_rows", "page_rows", "list_tables"}) + serverName := os.Getenv("VASTBASE_SERVER") + + encoder := json.NewEncoder(os.Stdout) + metadata := benchmarkMetadata{ + Type: "metadata", + Server: serverName, + GOOS: runtime.GOOS, + GOARCH: runtime.GOARCH, + Phases: phases, + Rounds: rounds, + DurationSeconds: durationSeconds, + Concurrencies: concurrencies, + StartupWarmups: startupWarmups, + StartupIterations: startupIterations, + ConnectWarmups: connectWarmups, + ConnectIterations: connectIterations, + QueryWarmups: queryWarmups, + Workloads: workloadNames, + } + encode(encoder, metadata) + + if contains(phases, "startup") { + for _, result := range benchmarkStartups(agents, startupWarmups, startupIterations) { + encode(encoder, result) + } + } + if !contains(phases, "connect") && !contains(phases, "query") { + return + } + + connection := connectionParams() + if serverName == "" { + serverName = fmt.Sprintf("%s:%v", connection["host"], connection["port"]) + } + maxConcurrency := maxInt(concurrencies) + running := startPersistentAgents(agents) + defer func() { + for _, candidate := range running { + _ = candidate.process.close() + } + }() + + versions := preflightVersions(running, connection) + metadata.Server = serverName + metadata.DatabaseVersions = versions + encode(encoder, metadata) + + if contains(phases, "connect") { + for _, result := range benchmarkConnections(running, connection, connectWarmups, connectIterations) { + result.Server = serverName + encode(encoder, result) + } + } + if !contains(phases, "query") { + return + } + + openSessions(running, connection, maxConcurrency) + workloads := configuredWorkloads(workloadNames) + for _, benchmark := range workloads { + for _, concurrency := range concurrencies { + for round := 1; round <= rounds; round++ { + for _, candidate := range rotatedAgents(running, round+concurrency) { + warmup(candidate.process, benchmark, concurrency, queryWarmups) + result := runWorkload( + candidate.process, + benchmark, + time.Duration(durationSeconds)*time.Second, + concurrency, + ) + result.Server = serverName + result.Agent = candidate.spec.Name + result.Round = round + result.ReadyRSSKB = candidate.readyRSSKB + result.OneSessionKB = candidate.oneSessionRSSKB + result.AllSessionsKB = candidate.allSessionsRSS + result.ArtifactBytes = fileSize(candidate.spec.ArtifactPath) + encode(encoder, result) + } + } + } + } +} + +func benchmarkStartups(agents []agentSpec, warmups, iterations int) []benchmarkResult { + for iteration := 0; iteration < warmups; iteration++ { + for _, agent := range rotatedSpecs(agents, iteration) { + process, _, err := startAgent(agent.Command) + if err != nil { + panic(fmt.Errorf("warm startup %s: %w", agent.Name, err)) + } + if _, err := process.call("handshake", map[string]any{}); err != nil { + process.kill() + panic(fmt.Errorf("warm handshake %s: %w", agent.Name, err)) + } + if err := process.close(); err != nil { + panic(fmt.Errorf("close startup warmup %s: %w", agent.Name, err)) + } + } + } + + readySamples := map[string][]float64{} + handshakeSamples := map[string][]float64{} + rssSamples := map[string][]int64{} + for iteration := 0; iteration < iterations; iteration++ { + for _, agent := range rotatedSpecs(agents, iteration) { + process, readyDuration, err := startAgent(agent.Command) + if err != nil { + panic(fmt.Errorf("start %s: %w", agent.Name, err)) + } + handshakeStart := time.Now() + if _, err := process.call("handshake", map[string]any{}); err != nil { + process.kill() + panic(fmt.Errorf("handshake %s: %w", agent.Name, err)) + } + readySamples[agent.Name] = append(readySamples[agent.Name], milliseconds(readyDuration)) + handshakeSamples[agent.Name] = append( + handshakeSamples[agent.Name], + milliseconds(readyDuration+time.Since(handshakeStart)), + ) + rssSamples[agent.Name] = append(rssSamples[agent.Name], readRSSKB(process.command.Process.Pid)) + if err := process.close(); err != nil { + panic(fmt.Errorf("close startup %s: %w", agent.Name, err)) + } + } + } + + results := make([]benchmarkResult, 0, len(agents)*2) + for _, agent := range agents { + ready := summarize(agent.Name, "startup_ready", 0, readySamples[agent.Name]) + ready.ReadyRSSKB = medianInt64(rssSamples[agent.Name]) + ready.ArtifactBytes = fileSize(agent.ArtifactPath) + results = append(results, ready) + + withHandshake := summarize(agent.Name, "startup_handshake", 0, handshakeSamples[agent.Name]) + withHandshake.ReadyRSSKB = medianInt64(rssSamples[agent.Name]) + withHandshake.ArtifactBytes = fileSize(agent.ArtifactPath) + results = append(results, withHandshake) + } + return results +} + +func startPersistentAgents(agents []agentSpec) []*runningAgent { + running := make([]*runningAgent, 0, len(agents)) + for _, agent := range agents { + process, _, err := startAgent(agent.Command) + if err != nil { + panic(fmt.Errorf("start persistent %s: %w", agent.Name, err)) + } + if _, err := process.call("handshake", map[string]any{}); err != nil { + process.kill() + panic(fmt.Errorf("handshake persistent %s: %w", agent.Name, err)) + } + running = append(running, &runningAgent{ + spec: agent, + process: process, + readyRSSKB: readRSSKB(process.command.Process.Pid), + }) + } + return running +} + +func preflightVersions(running []*runningAgent, connection map[string]any) map[string]string { + versions := map[string]string{} + for _, candidate := range running { + params := cloneMap(connection) + params["agentSessionId"] = "preflight" + if _, err := candidate.process.call("open_session", params); err != nil { + panic(fmt.Errorf("preflight connect %s: %w", candidate.spec.Name, err)) + } + result, err := candidate.process.call("execute_query", map[string]any{ + "agentSessionId": "preflight", + "sql": "SELECT version()", + "maxRows": 1, + }) + if err != nil { + panic(fmt.Errorf("preflight version %s: %w", candidate.spec.Name, err)) + } + versions[candidate.spec.Name] = firstCell(result) + if _, err := candidate.process.call("close_session", map[string]any{"agentSessionId": "preflight"}); err != nil { + panic(fmt.Errorf("close preflight %s: %w", candidate.spec.Name, err)) + } + } + return versions +} + +func benchmarkConnections( + running []*runningAgent, + connection map[string]any, + warmups int, + iterations int, +) []benchmarkResult { + for iteration := 0; iteration < warmups; iteration++ { + for _, candidate := range rotatedAgents(running, iteration) { + benchmarkOneConnection(candidate, connection, fmt.Sprintf("connect-warmup-%d", iteration)) + } + } + + samples := map[string][]float64{} + for iteration := 0; iteration < iterations; iteration++ { + for _, candidate := range rotatedAgents(running, iteration) { + start := time.Now() + benchmarkOneConnection(candidate, connection, fmt.Sprintf("connect-%d", iteration)) + samples[candidate.spec.Name] = append(samples[candidate.spec.Name], milliseconds(time.Since(start))) + } + } + + results := make([]benchmarkResult, 0, len(running)) + for _, candidate := range running { + result := summarize(candidate.spec.Name, "connect_auth_close", 1, samples[candidate.spec.Name]) + result.ReadyRSSKB = readRSSKB(candidate.process.command.Process.Pid) + result.ArtifactBytes = fileSize(candidate.spec.ArtifactPath) + results = append(results, result) + } + return results +} + +func benchmarkOneConnection(candidate *runningAgent, connection map[string]any, session string) { + params := cloneMap(connection) + params["agentSessionId"] = session + if _, err := candidate.process.call("open_session", params); err != nil { + panic(fmt.Errorf("open connection %s: %w", candidate.spec.Name, err)) + } + if _, err := candidate.process.call("close_session", map[string]any{"agentSessionId": session}); err != nil { + panic(fmt.Errorf("close connection %s: %w", candidate.spec.Name, err)) + } +} + +func openSessions(running []*runningAgent, connection map[string]any, count int) { + for _, candidate := range running { + for index := 0; index < count; index++ { + params := cloneMap(connection) + params["agentSessionId"] = sessionID(index) + if _, err := candidate.process.call("open_session", params); err != nil { + panic(fmt.Errorf("open %s session %d: %w", candidate.spec.Name, index, err)) + } + if index == 0 { + candidate.oneSessionRSSKB = readRSSKB(candidate.process.command.Process.Pid) + } + } + candidate.allSessionsRSS = readRSSKB(candidate.process.command.Process.Pid) + } +} + +func configuredWorkloads(names []string) []workload { + literalSQL := envOr("BENCH_LITERAL_SQL", "SELECT 1 AS value") + decodeRows := envInt("BENCH_DECODE_ROWS", 1000) + decodeSQL := envOr( + "BENCH_DECODE_SQL", + fmt.Sprintf( + "SELECT value AS id, CAST(value * 1.25 AS numeric(18,2)) AS numeric_value, "+ + "CAST('2024-01-02 03:04:05' AS timestamp) AS timestamp_value, repeat('x', 64) AS text_value "+ + "FROM generate_series(1, %d) AS value", + decodeRows, + ), + ) + pageRows := envInt("BENCH_PAGE_ROWS", decodeRows) + pageSQL := envOr("BENCH_PAGE_SQL", decodeSQL) + schema := envOr("BENCH_SCHEMA", "public") + + available := map[string]workload{ + "select_literal": { + Name: "select_literal", + Method: "execute_query", + Parameters: func(worker int) map[string]any { + return map[string]any{ + "agentSessionId": sessionID(worker), + "sql": literalSQL, + "maxRows": 1, + } + }, + }, + "decode_rows": { + Name: "decode_rows", + Method: "execute_query", + Parameters: func(worker int) map[string]any { + return map[string]any{ + "agentSessionId": sessionID(worker), + "sql": decodeSQL, + "maxRows": decodeRows, + "fetchSize": decodeRows, + } + }, + }, + "page_rows": { + Name: "page_rows", + Method: "execute_query_page", + Parameters: func(worker int) map[string]any { + return map[string]any{ + "agentSessionId": sessionID(worker), + "sql": pageSQL, + "pageSize": pageRows, + "fetchSize": pageRows, + "maxRows": pageRows, + } + }, + Cleanup: cleanupQueryPage, + }, + "list_tables": { + Name: "list_tables", + Method: "list_tables", + Parameters: func(worker int) map[string]any { + return map[string]any{"agentSessionId": sessionID(worker), "schema": schema} + }, + }, + } + + result := make([]workload, 0, len(names)) + for _, name := range names { + benchmark, ok := available[name] + if !ok { + panic("unknown BENCH_WORKLOADS entry: " + name) + } + result = append(result, benchmark) + } + return result +} + +func cleanupQueryPage(process *agentProcess, result json.RawMessage, worker int) error { + var page struct { + SessionID string `json:"sessionId"` + Done bool `json:"done"` + } + if err := json.Unmarshal(result, &page); err != nil || page.Done || page.SessionID == "" { + return err + } + _, err := process.call("close_query_session", map[string]any{ + "agentSessionId": sessionID(worker), + "sessionId": page.SessionID, + }) + return err +} + +func warmup(process *agentProcess, benchmark workload, concurrency, operations int) { + if operations == 0 { + return + } + for iteration := 0; iteration < operations; iteration++ { + worker := iteration % concurrency + result, err := process.call(benchmark.Method, benchmark.Parameters(worker)) + if err != nil { + panic(fmt.Errorf("warmup %s: %w", benchmark.Name, err)) + } + if benchmark.Cleanup != nil { + if err := benchmark.Cleanup(process, result, worker); err != nil { + panic(fmt.Errorf("warmup cleanup %s: %w", benchmark.Name, err)) + } + } + } +} + +func runWorkload(process *agentProcess, benchmark workload, duration time.Duration, concurrency int) benchmarkResult { + var operations atomic.Int64 + var failures atomic.Int64 + var peakRSS atomic.Int64 + peakRSS.Store(readRSSKB(process.command.Process.Pid)) + stopMemory := make(chan struct{}) + go func() { + ticker := time.NewTicker(100 * time.Millisecond) + defer ticker.Stop() + for { + select { + case <-ticker.C: + value := readRSSKB(process.command.Process.Pid) + for value > peakRSS.Load() && !peakRSS.CompareAndSwap(peakRSS.Load(), value) { + } + case <-stopMemory: + return + } + } + }() + + latencies := make([][]float64, concurrency) + start := time.Now() + deadline := start.Add(duration) + var workers sync.WaitGroup + for worker := 0; worker < concurrency; worker++ { + worker := worker + workers.Add(1) + go func() { + defer workers.Done() + local := make([]float64, 0, 4096) + for time.Now().Before(deadline) { + callStart := time.Now() + result, err := process.call(benchmark.Method, benchmark.Parameters(worker)) + if err == nil && benchmark.Cleanup != nil { + err = benchmark.Cleanup(process, result, worker) + } + local = append(local, milliseconds(time.Since(callStart))) + operations.Add(1) + if err != nil { + failures.Add(1) + } + } + latencies[worker] = local + }() + } + workers.Wait() + close(stopMemory) + elapsed := time.Since(start) + + merged := make([]float64, 0) + for _, values := range latencies { + merged = append(merged, values...) + } + result := summarize("", benchmark.Name, concurrency, merged) + result.DurationMS = milliseconds(elapsed) + result.QPS = float64(operations.Load()) / elapsed.Seconds() + result.Operations = operations.Load() + result.Errors = failures.Load() + result.PeakRSSKB = peakRSS.Load() + return result +} + +func startAgent(argv []string) (*agentProcess, time.Duration, error) { + if len(argv) == 0 { + return nil, 0, errors.New("agent command is empty") + } + command := exec.Command(argv[0], argv[1:]...) + stdin, err := command.StdinPipe() + if err != nil { + return nil, 0, err + } + stdout, err := command.StdoutPipe() + if err != nil { + return nil, 0, err + } + command.Stderr = os.Stderr + process := &agentProcess{command: command, stdin: stdin, reader: bufio.NewScanner(stdout)} + process.reader.Buffer(make([]byte, 0, 64*1024), 512*1024*1024) + start := time.Now() + if err := command.Start(); err != nil { + return nil, 0, err + } + if !process.reader.Scan() { + return nil, 0, errors.New("agent exited before ready") + } + if !strings.Contains(process.reader.Text(), `"ready":true`) { + process.kill() + return nil, 0, fmt.Errorf("agent did not become ready: %s", process.reader.Text()) + } + readyDuration := time.Since(start) + go process.readResponses() + return process, readyDuration, nil +} + +func (process *agentProcess) readResponses() { + for process.reader.Scan() { + var response agentResponse + if json.Unmarshal(process.reader.Bytes(), &response) != nil { + continue + } + if channel, ok := process.pending.LoadAndDelete(response.ID); ok { + channel.(chan agentResponse) <- response + } + } +} + +func (process *agentProcess) call(method string, params map[string]any) (json.RawMessage, error) { + id := process.nextID.Add(1) + channel := make(chan agentResponse, 1) + process.pending.Store(id, channel) + request := map[string]any{"id": id, "method": method, "params": params} + payload, err := json.Marshal(request) + if err != nil { + process.pending.Delete(id) + return nil, err + } + process.writeMu.Lock() + _, err = process.stdin.Write(append(payload, '\n')) + process.writeMu.Unlock() + if err != nil { + process.pending.Delete(id) + return nil, err + } + select { + case response := <-channel: + if response.Error != nil { + return nil, errors.New(response.Error.Message) + } + return response.Result, nil + case <-time.After(60 * time.Second): + process.pending.Delete(id) + return nil, errors.New("agent request timed out") + } +} + +func (process *agentProcess) close() error { + _, _ = process.call("shutdown", map[string]any{}) + _ = process.stdin.Close() + return process.command.Wait() +} + +func (process *agentProcess) kill() { + if process.command.Process != nil { + _ = process.command.Process.Kill() + } +} + +func connectionParams() map[string]any { + port, err := strconv.Atoi(requiredEnv("VASTBASE_PORT")) + if err != nil { + panic(fmt.Errorf("parse VASTBASE_PORT: %w", err)) + } + return map[string]any{ + "host": requiredEnv("VASTBASE_HOST"), + "port": port, + "database": requiredEnv("VASTBASE_DATABASE"), + "username": requiredEnv("VASTBASE_USERNAME"), + "password": requiredEnv("VASTBASE_PASSWORD"), + "url_params": os.Getenv("VASTBASE_URL_PARAMS"), + "connection_string": os.Getenv("VASTBASE_CONNECTION_STRING"), + "ssl": envBool("VASTBASE_SSL", false), + "ca_cert_path": os.Getenv("VASTBASE_CA_CERT_PATH"), + "client_cert_path": os.Getenv("VASTBASE_CLIENT_CERT_PATH"), + "client_key_path": os.Getenv("VASTBASE_CLIENT_KEY_PATH"), + } +} + +func jdbcAgentCommand(jar string) []string { + java := os.Getenv("DBX_AGENT_JAVA") + if java == "" { + java = "java" + } + return []string{java, "-Xms32m", "-Xmx512m", "-jar", jar} +} + +func summarize(agent, workload string, concurrency int, values []float64) benchmarkResult { + sorted := append([]float64(nil), values...) + sort.Float64s(sorted) + var total float64 + for _, value := range sorted { + total += value + } + durationMS := total + qps := 0.0 + if durationMS > 0 { + qps = float64(len(sorted)) / (durationMS / 1000) + } + return benchmarkResult{ + Type: "result", + Agent: agent, + Workload: workload, + Concurrency: concurrency, + Operations: int64(len(sorted)), + DurationMS: durationMS, + QPS: qps, + MeanMS: total / float64(maxInt([]int{1, len(sorted)})), + P50MS: percentile(sorted, 0.50), + P95MS: percentile(sorted, 0.95), + P99MS: percentile(sorted, 0.99), + } +} + +func percentile(values []float64, fraction float64) float64 { + if len(values) == 0 { + return 0 + } + index := int(float64(len(values)-1) * fraction) + return values[index] +} + +func firstCell(result json.RawMessage) string { + var query struct { + Rows [][]any `json:"rows"` + } + if json.Unmarshal(result, &query) != nil || len(query.Rows) == 0 || len(query.Rows[0]) == 0 { + return "" + } + return fmt.Sprint(query.Rows[0][0]) +} + +func rotatedSpecs(values []agentSpec, offset int) []agentSpec { + if len(values) == 0 { + return nil + } + start := offset % len(values) + result := make([]agentSpec, 0, len(values)) + result = append(result, values[start:]...) + result = append(result, values[:start]...) + return result +} + +func rotatedAgents(values []*runningAgent, offset int) []*runningAgent { + if len(values) == 0 { + return nil + } + start := offset % len(values) + result := make([]*runningAgent, 0, len(values)) + result = append(result, values[start:]...) + result = append(result, values[:start]...) + return result +} + +func readRSSKB(pid int) int64 { + output, err := exec.Command("ps", "-o", "rss=", "-p", strconv.Itoa(pid)).Output() + if err != nil { + return 0 + } + value, _ := strconv.ParseInt(strings.TrimSpace(string(output)), 10, 64) + return value +} + +func medianInt64(values []int64) int64 { + if len(values) == 0 { + return 0 + } + sorted := append([]int64(nil), values...) + sort.Slice(sorted, func(left, right int) bool { return sorted[left] < sorted[right] }) + return sorted[len(sorted)/2] +} + +func fileSize(path string) int64 { + info, err := os.Stat(path) + if err != nil { + return 0 + } + return info.Size() +} + +func cloneMap(source map[string]any) map[string]any { + result := make(map[string]any, len(source)+1) + for key, value := range source { + result[key] = value + } + return result +} + +func sessionID(index int) string { + return "bench-" + strconv.Itoa(index) +} + +func encode(encoder *json.Encoder, value any) { + if err := encoder.Encode(value); err != nil { + panic(err) + } +} + +func milliseconds(value time.Duration) float64 { + return float64(value.Microseconds()) / 1000 +} + +func contains(values []string, expected string) bool { + for _, value := range values { + if value == expected { + return true + } + } + return false +} + +func maxInt(values []int) int { + result := 0 + for _, value := range values { + if value > result { + result = value + } + } + return result +} + +func requiredEnv(name string) string { + value := os.Getenv(name) + if value == "" { + panic(name + " is required") + } + return value +} + +func envOr(name, fallback string) string { + if value := os.Getenv(name); value != "" { + return value + } + return fallback +} + +func envInt(name string, fallback int) int { + value, err := strconv.Atoi(os.Getenv(name)) + if err != nil || value <= 0 { + return fallback + } + return value +} + +func envNonNegativeInt(name string, fallback int) int { + value, err := strconv.Atoi(os.Getenv(name)) + if err != nil || value < 0 { + return fallback + } + return value +} + +func envBool(name string, fallback bool) bool { + value := strings.TrimSpace(os.Getenv(name)) + if value == "" { + return fallback + } + parsed, err := strconv.ParseBool(value) + if err != nil { + panic(fmt.Errorf("parse %s: %w", name, err)) + } + return parsed +} + +func envStrings(name string, fallback []string) []string { + raw := strings.TrimSpace(os.Getenv(name)) + if raw == "" { + return fallback + } + result := make([]string, 0) + for _, item := range strings.Split(raw, ",") { + if value := strings.TrimSpace(item); value != "" { + result = append(result, value) + } + } + if len(result) == 0 { + return fallback + } + return result +} + +func envInts(name string, fallback []int) []int { + items := envStrings(name, nil) + if len(items) == 0 { + return fallback + } + result := make([]int, 0, len(items)) + for _, item := range items { + value, err := strconv.Atoi(item) + if err != nil || value <= 0 { + panic(name + " must contain positive integers") + } + result = append(result, value) + } + return result +} diff --git a/agents/drivers/vastbase-go/bench/direct/main.go b/agents/drivers/vastbase-go/bench/direct/main.go new file mode 100644 index 000000000..ed369cf98 --- /dev/null +++ b/agents/drivers/vastbase-go/bench/direct/main.go @@ -0,0 +1,209 @@ +package main + +import ( + "context" + "database/sql" + "encoding/base64" + "encoding/json" + "fmt" + "os" + "strconv" + "sync" + "sync/atomic" + "time" + + _ "gitcode.com/opengauss/openGauss-connector-go-pq" +) + +type queryResult struct { + Columns []string `json:"columns"` + ColumnTypes []string `json:"column_types"` + Rows [][]any `json:"rows"` +} + +type benchmarkResult struct { + Mode string `json:"mode"` + Concurrency int `json:"concurrency"` + Seconds int `json:"seconds"` + Operations int64 `json:"operations"` + Errors int64 `json:"errors"` + QPS float64 `json:"qps"` +} + +func main() { + concurrency := envInt("BENCH_CONCURRENCY", 32) + seconds := envInt("BENCH_SECONDS", 4) + mode := envOr("BENCH_MODE", "collect") + db, err := sql.Open("opengauss", dsn()) + if err != nil { + panic(err) + } + defer db.Close() + db.SetMaxOpenConns(concurrency) + db.SetMaxIdleConns(concurrency) + + connections := make([]*sql.Conn, concurrency) + for index := range connections { + connections[index], err = db.Conn(context.Background()) + if err != nil { + panic(err) + } + defer connections[index].Close() + } + for _, connection := range connections { + for range 2 { + if err := execute(connection, mode); err != nil { + panic(err) + } + } + } + + start := time.Now() + deadline := start.Add(time.Duration(seconds) * time.Second) + var operations atomic.Int64 + var failures atomic.Int64 + var waitGroup sync.WaitGroup + for _, connection := range connections { + waitGroup.Add(1) + go func(connection *sql.Conn) { + defer waitGroup.Done() + for time.Now().Before(deadline) { + if err := execute(connection, mode); err != nil { + failures.Add(1) + } else { + operations.Add(1) + } + } + }(connection) + } + waitGroup.Wait() + duration := time.Since(start).Seconds() + result := benchmarkResult{ + Mode: mode, + Concurrency: concurrency, + Seconds: seconds, + Operations: operations.Load(), + Errors: failures.Load(), + QPS: float64(operations.Load()) / duration, + } + if err := json.NewEncoder(os.Stdout).Encode(result); err != nil { + panic(err) + } +} + +func execute(connection *sql.Conn, mode string) error { + rows, err := connection.QueryContext(context.Background(), querySQL()) + if err != nil { + return err + } + defer rows.Close() + columns, err := rows.Columns() + if err != nil { + return err + } + types, err := rows.ColumnTypes() + if err != nil { + return err + } + columnTypes := make([]string, len(types)) + for index, columnType := range types { + columnTypes[index] = columnType.DatabaseTypeName() + } + values := make([]any, len(columns)) + destinations := make([]any, len(columns)) + for index := range values { + destinations[index] = &values[index] + } + result := queryResult{Columns: columns, ColumnTypes: columnTypes, Rows: make([][]any, 0, 1000)} + for rows.Next() { + if err := rows.Scan(destinations...); err != nil { + return err + } + row := make([]any, len(values)) + for index, value := range values { + row[index] = normalizeValue(value) + } + result.Rows = append(result.Rows, row) + } + if err := rows.Err(); err != nil { + return err + } + if mode == "marshal" { + _, err = json.Marshal(result) + return err + } + return nil +} + +func normalizeValue(value any) any { + switch typed := value.(type) { + case nil: + return nil + case []byte: + if isTextBytes(typed) { + return string(typed) + } + return map[string]string{"$binary": base64.StdEncoding.EncodeToString(typed)} + case time.Time: + return typed.Format(time.RFC3339Nano) + case int8: + return int64(typed) + case int16: + return int64(typed) + case int32: + return int64(typed) + case float32: + return float64(typed) + default: + return typed + } +} + +func isTextBytes(value []byte) bool { + for _, char := range value { + if char == 0 || char < 0x09 || char > 0x0d && char < 0x20 { + return false + } + } + return true +} + +func dsn() string { + return fmt.Sprintf( + "host=%s port=%s user=%s password=%s dbname=%s sslmode=disable", + envOr("VASTBASE_HOST", "127.0.0.1"), + envOr("VASTBASE_PORT", "20119"), + envOr("VASTBASE_USERNAME", "dbx_bench"), + requiredEnv("DBX_TEST_PASSWORD"), + envOr("VASTBASE_DATABASE", "dbx_bench"), + ) +} + +func querySQL() string { + return "SELECT value AS id, CAST(value * 1.25 AS numeric(18,2)) AS numeric_value, " + + "CAST('2024-01-02 03:04:05' AS timestamp) AS timestamp_value, repeat('x', 64) AS text_value " + + "FROM generate_series(1, 1000) AS value" +} + +func envInt(key string, fallback int) int { + value, err := strconv.Atoi(os.Getenv(key)) + if err == nil && value > 0 { + return value + } + return fallback +} + +func envOr(key, fallback string) string { + if value := os.Getenv(key); value != "" { + return value + } + return fallback +} + +func requiredEnv(key string) string { + value := os.Getenv(key) + if value == "" { + panic(key + " is required") + } + return value +} diff --git a/agents/drivers/vastbase-go/connection_info_test.go b/agents/drivers/vastbase-go/connection_info_test.go new file mode 100644 index 000000000..972b11fbf --- /dev/null +++ b/agents/drivers/vastbase-go/connection_info_test.go @@ -0,0 +1,112 @@ +package main + +import ( + "context" + "database/sql" + "database/sql/driver" + "fmt" + "io" + "sync/atomic" + "testing" +) + +func TestConnectionInfoAndTestConnectionExposeDatabaseInfo(t *testing.T) { + db := openConnectionInfoTestDB(t) + server := newServer() + server.db = db + server.mode = detectAgentMode(db, false) + + info, err := server.connectionInfo() + if err != nil { + t.Fatal(err) + } + assertVastbaseDatabaseInfo(t, info["databaseInfo"]) + + testServer := newServer() + testServer.openDatabase = func(connectParams, string) (*sql.DB, error) { + return openConnectionInfoTestDB(t), nil + } + result, err := testServer.testConnection(connectParams{}) + if err != nil { + t.Fatal(err) + } + if ok, _ := result["ok"].(bool); !ok { + t.Fatalf("test_connection did not succeed: %v", result) + } + assertVastbaseDatabaseInfo(t, result["databaseInfo"]) +} + +func assertVastbaseDatabaseInfo(t *testing.T, value any) { + t.Helper() + info, ok := value.(map[string]string) + if !ok { + t.Fatalf("unexpected databaseInfo type: %T", value) + } + expected := map[string]string{ + "productName": "Vastbase", + "productVersion": "Vastbase G100 V3.0.9", + "unquotedIdentifierCase": "lower", + "quotedIdentifierCase": "mixed", + "driverName": agentDriverName, + "driverVersion": agentDriverVersion, + } + for key, expectedValue := range expected { + if info[key] != expectedValue { + t.Fatalf("databaseInfo[%s] = %q, want %q", key, info[key], expectedValue) + } + } +} + +var connectionInfoDriverSequence atomic.Uint64 + +type connectionInfoTestDriver struct{} + +func (*connectionInfoTestDriver) Open(string) (driver.Conn, error) { + return &connectionInfoTestConn{}, nil +} + +type connectionInfoTestConn struct{} + +func (*connectionInfoTestConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip } +func (*connectionInfoTestConn) Close() error { return nil } +func (*connectionInfoTestConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip } +func (*connectionInfoTestConn) Ping(context.Context) error { return nil } + +func (*connectionInfoTestConn) QueryContext(_ context.Context, _ string, _ []driver.NamedValue) (driver.Rows, error) { + return &connectionInfoTestRows{ + columns: []string{"current_database", "current_user", "version", "current_schema"}, + values: []driver.Value{"postgres", "vbadmin", "Vastbase G100 V3.0.9", "public"}, + }, nil +} + +type connectionInfoTestRows struct { + columns []string + values []driver.Value + done bool +} + +func (rows *connectionInfoTestRows) Columns() []string { return rows.columns } +func (*connectionInfoTestRows) Close() error { return nil } + +func (rows *connectionInfoTestRows) Next(destination []driver.Value) error { + if rows.done { + return io.EOF + } + copy(destination, rows.values) + rows.done = true + return nil +} + +func openConnectionInfoTestDB(t *testing.T) *sql.DB { + t.Helper() + driverName := fmt.Sprintf("vastbase-connection-info-%d", connectionInfoDriverSequence.Add(1)) + sql.Register(driverName, &connectionInfoTestDriver{}) + db, err := sql.Open(driverName, "") + if err != nil { + t.Fatal(err) + } + db.SetMaxOpenConns(1) + db.SetMaxIdleConns(1) + t.Cleanup(func() { _ = db.Close() }) + return db +} diff --git a/agents/drivers/vastbase-go/connection_state.go b/agents/drivers/vastbase-go/connection_state.go new file mode 100644 index 000000000..1569c53b2 --- /dev/null +++ b/agents/drivers/vastbase-go/connection_state.go @@ -0,0 +1,222 @@ +package main + +import ( + "database/sql" + "reflect" + "regexp" + "strings" +) + +var ( + sessionAffinityFunction = regexp.MustCompile(`(?is)\b(?:SET_CONFIG|PG_(?:TRY_)?ADVISORY_(?:(?:XACT_)?LOCK(?:_SHARED)?|UNLOCK(?:_SHARED|_ALL)?)|GET_LOCK|RELEASE_LOCK|SP_GETAPPLOCK|DBMS_LOCK)\s*\(`) + sessionUserVariable = regexp.MustCompile(`(?is)(?:SET\s+)?@[A-Z0-9_$]+\s*(?::=|=)`) + sessionTemporaryObject = regexp.MustCompile(`(?is)(?:^|[^A-Z0-9_$])#{1,2}[A-Z0-9_$]+`) +) + +func sqlConnectionIdentity(conn *sql.Conn) uintptr { + var identity uintptr + _ = conn.Raw(func(raw any) error { + value := reflect.ValueOf(raw) + if value.IsValid() && value.Kind() == reflect.Pointer { + identity = value.Pointer() + } + return nil + }) + return identity +} + +func (s *server) resetSchemaCache() { + s.currentSchema = "" + s.schemaInitialized = false + s.schemaConnectionID = 0 +} + +func (s *server) invalidateSchemaAfterSQL(sqlText string) { + if sqlMayChangeSessionState(sqlText) { + s.resetSchemaCache() + } +} + +func (s *server) noteSQLSessionState(sqlText string) { + s.invalidateSchemaAfterSQL(sqlText) + if sqlRequiresSessionAffinity(sqlText) { + s.sessionAffinity = true + } +} + +func sqlRequiresSessionAffinity(sqlText string) bool { + normalized := strings.ToUpper(sanitizeSessionStateSQL(sqlText)) + if sessionAffinityFunction.MatchString(normalized) || sessionUserVariable.MatchString(normalized) || sessionTemporaryObject.MatchString(normalized) { + return true + } + for _, statement := range strings.Split(normalized, ";") { + fields := strings.Fields(statement) + if len(fields) == 0 { + continue + } + switch fields[0] { + case "BEGIN", "SET", "RESET", "UNSET", "USE", "DATABASE", "DECLARE", "PREPARE", "DEALLOCATE", "ATTACH", "DETACH", "PRAGMA", "CALL", "EXEC", "EXECUTE", "DO", "LISTEN", "UNLISTEN", "LOAD", "INSTALL": + return true + case "START": + if len(fields) > 1 && fields[1] == "TRANSACTION" { + return true + } + case "ALTER": + if len(fields) > 1 && fields[1] == "SESSION" { + return true + } + case "LOCK", "UNLOCK": + if len(fields) > 1 && strings.HasPrefix(fields[1], "TABLE") { + return true + } + case "CREATE": + for _, field := range fields[1:] { + if field == "TEMP" || field == "TEMPORARY" || field == "VOLATILE" { + return true + } + if field == "TABLE" { + break + } + } + case "SELECT": + for index, field := range fields { + if field == "INTO" && index+1 < len(fields) && (fields[index+1] == "TEMP" || fields[index+1] == "TEMPORARY") { + return true + } + } + case "ADD", "DELETE": + if len(fields) > 1 && (fields[1] == "JAR" || fields[1] == "FILE" || fields[1] == "ARCHIVE") { + return true + } + case "CACHE", "UNCACHE": + if len(fields) > 1 && fields[1] == "TABLE" { + return true + } + } + } + return false +} + +func sqlMayChangeSessionState(sqlText string) bool { + normalized := strings.ToUpper(sanitizeSessionStateSQL(sqlText)) + if strings.Contains(normalized, "SET_CONFIG") { + return true + } + for _, statement := range strings.Split(normalized, ";") { + fields := strings.Fields(statement) + if len(fields) == 0 { + continue + } + switch fields[0] { + case "SET", "RESET", "DISCARD": + return true + case "ALTER": + if len(fields) > 1 && fields[1] == "SESSION" { + return true + } + } + } + return false +} + +func sanitizeSessionStateSQL(sqlText string) string { + var sanitized strings.Builder + sanitized.Grow(len(sqlText)) + for index := 0; index < len(sqlText); { + switch { + case index+1 < len(sqlText) && sqlText[index] == '-' && sqlText[index+1] == '-': + index = sanitizeSQLLine(sqlText, &sanitized, index, index+2) + case sqlText[index] == '#': + index = sanitizeSQLLine(sqlText, &sanitized, index, index+1) + case index+1 < len(sqlText) && sqlText[index] == '/' && sqlText[index+1] == '*': + index = sanitizeSQLBlock(sqlText, &sanitized, index+2) + case sqlText[index] == '\'' || sqlText[index] == '"' || sqlText[index] == '`': + index = sanitizeSQLQuoted(sqlText, &sanitized, index, sqlText[index]) + case sqlText[index] == '[': + index = sanitizeSQLQuoted(sqlText, &sanitized, index, ']') + case sqlText[index] == '$': + delimiter := sqlDollarQuoteDelimiter(sqlText, index) + if delimiter == "" { + sanitized.WriteByte(sqlText[index]) + index++ + continue + } + closing := strings.Index(sqlText[index+len(delimiter):], delimiter) + if closing < 0 { + sanitized.WriteByte(sqlText[index]) + index++ + continue + } + end := index + len(delimiter) + closing + len(delimiter) + appendSanitizedSQL(sqlText, &sanitized, index, end) + index = end + default: + sanitized.WriteByte(sqlText[index]) + index++ + } + } + return sanitized.String() +} + +func sanitizeSQLLine(sqlText string, sanitized *strings.Builder, start, index int) int { + for index < len(sqlText) && sqlText[index] != '\n' && sqlText[index] != '\r' { + index++ + } + appendSanitizedSQL(sqlText, sanitized, start, index) + return index +} + +func sanitizeSQLBlock(sqlText string, sanitized *strings.Builder, index int) int { + start := index - 2 + closing := strings.Index(sqlText[index:], "*/") + end := len(sqlText) + if closing >= 0 { + end = index + closing + 2 + } + appendSanitizedSQL(sqlText, sanitized, start, end) + return end +} + +func sanitizeSQLQuoted(sqlText string, sanitized *strings.Builder, start int, closing byte) int { + index := start + 1 + for index < len(sqlText) { + if sqlText[index] == closing { + if index+1 < len(sqlText) && sqlText[index+1] == closing { + index += 2 + continue + } + index++ + break + } + if sqlText[index] == '\\' && index+1 < len(sqlText) { + index += 2 + continue + } + index++ + } + appendSanitizedSQL(sqlText, sanitized, start, index) + return index +} + +func sqlDollarQuoteDelimiter(sqlText string, start int) string { + for index := start + 1; index < len(sqlText); index++ { + if sqlText[index] == '$' { + return sqlText[start : index+1] + } + char := sqlText[index] + if !((char >= 'a' && char <= 'z') || (char >= 'A' && char <= 'Z') || (char >= '0' && char <= '9') || char == '_') { + return "" + } + } + return "" +} + +func appendSanitizedSQL(sqlText string, sanitized *strings.Builder, start, end int) { + for index := start; index < end; index++ { + if sqlText[index] == '\n' || sqlText[index] == '\r' || sqlText[index] == ';' { + sanitized.WriteByte(sqlText[index]) + } else { + sanitized.WriteByte(' ') + } + } +} diff --git a/agents/drivers/vastbase-go/connection_state_test.go b/agents/drivers/vastbase-go/connection_state_test.go new file mode 100644 index 000000000..4d582322c --- /dev/null +++ b/agents/drivers/vastbase-go/connection_state_test.go @@ -0,0 +1,231 @@ +package main + +import ( + "context" + "database/sql" + "database/sql/driver" + "fmt" + "sync" + "sync/atomic" + "testing" +) + +func TestSchemaCacheTracksPhysicalConnectionAndSchema(t *testing.T) { + state := &schemaCacheTestState{invalid: map[int]bool{}} + db := openSchemaCacheTestDB(t, state) + server := newServer() + server.db = db + + conn := mustSchemaConn(t, server, "") + _ = conn.Close() + if statements := state.executedStatements(); len(statements) != 0 { + t.Fatalf("fresh connection should preserve its initial schema: %v", statements) + } + + conn = mustSchemaConn(t, server, "public") + _ = conn.Close() + conn = mustSchemaConn(t, server, "public") + _ = conn.Close() + conn = mustSchemaConn(t, server, "analytics") + _ = conn.Close() + conn = mustSchemaConn(t, server, "") + _ = conn.Close() + + expected := []string{`SET search_path TO "public"`, `SET search_path TO "analytics"`, "RESET search_path"} + if statements := state.executedStatements(); fmt.Sprint(statements) != fmt.Sprint(expected) { + t.Fatalf("unexpected schema statements: got %v want %v", statements, expected) + } +} + +func TestSchemaCacheReappliesSchemaAfterPhysicalConnectionReplacement(t *testing.T) { + state := &schemaCacheTestState{invalid: map[int]bool{}} + db := openSchemaCacheTestDB(t, state) + server := newServer() + server.db = db + + conn := mustSchemaConn(t, server, "public") + state.invalidateLatest() + _ = conn.Close() + conn = mustSchemaConn(t, server, "public") + _ = conn.Close() + + if opens := state.openCount(); opens != 2 { + t.Fatalf("expected replacement physical connection, opened %d", opens) + } + expected := []string{`SET search_path TO "public"`, `SET search_path TO "public"`} + if statements := state.executedStatements(); fmt.Sprint(statements) != fmt.Sprint(expected) { + t.Fatalf("schema was not reapplied after replacement: got %v want %v", statements, expected) + } +} + +func TestSessionStateSQLInvalidatesSchemaCache(t *testing.T) { + mutating := []string{ + "SET search_path TO app", + "SELECT 1; \n SET ROLE analyst", + "RESET ALL", + "DISCARD ALL", + "ALTER SESSION SET CURRENT_SCHEMA = app", + "SELECT set_config('search_path', 'app', false)", + "SELECT pg_catalog.set_config ('role', 'analyst', false)", + "-- switch schema\nSET search_path TO app", + "SELECT 1; /* switch role */ RESET ROLE", + } + for _, statement := range mutating { + server := newServer() + server.currentSchema = "public" + server.schemaInitialized = true + server.schemaConnectionID = 1 + server.invalidateSchemaAfterSQL(statement) + if server.schemaInitialized { + t.Fatalf("schema cache was not invalidated for %q", statement) + } + } + + for _, statement := range []string{ + "SELECT 1", + "SELECT 'SET search_path TO app'", + "SELECT 'set_config(''search_path'', ''app'', false)'", + "SELECT $$RESET ALL$$", + "ALTER TABLE t ADD COLUMN value integer", + } { + server := newServer() + server.currentSchema = "public" + server.schemaInitialized = true + server.schemaConnectionID = 1 + server.invalidateSchemaAfterSQL(statement) + if !server.schemaInitialized { + t.Fatalf("schema cache was unnecessarily invalidated for %q", statement) + } + } +} + +func TestSQLRequiresSessionAffinity(t *testing.T) { + for _, statement := range []string{ + "CREATE TEMP TABLE scratch(id integer)", + "CREATE TEMPORARY TABLE scratch(id integer)", + "SELECT 1 INTO TEMP scratch", + "BEGIN", + "START TRANSACTION", + "SET ROLE analyst", + "RESET ROLE", + "SELECT pg_advisory_lock(42)", + "SELECT pg_try_advisory_xact_lock(42)", + "SELECT pg_advisory_unlock_all()", + "SELECT set_config('search_path', 'app', false)", + } { + if !sqlRequiresSessionAffinity(statement) { + t.Fatalf("session affinity was not detected for %q", statement) + } + } + + for _, statement := range []string{ + "SELECT 1", + "CREATE TABLE durable(id integer)", + "SELECT 'BEGIN; SET ROLE analyst'", + "SELECT $$CREATE TEMP TABLE scratch(id integer)$$", + "SELECT \"pg_advisory_lock\" FROM functions", + "-- SET ROLE analyst\nSELECT 1", + "/* CREATE TEMP TABLE scratch(id integer) */ SELECT 1", + } { + if sqlRequiresSessionAffinity(statement) { + t.Fatalf("session affinity was falsely detected for %q", statement) + } + } +} + +var schemaCacheDriverSequence atomic.Uint64 + +type schemaCacheTestState struct { + mu sync.Mutex + opens int + latestID int + invalid map[int]bool + statements []string +} + +func (state *schemaCacheTestState) openConnection() *schemaCacheTestConn { + state.mu.Lock() + defer state.mu.Unlock() + state.opens++ + state.latestID = state.opens + return &schemaCacheTestConn{state: state, id: state.latestID} +} + +func (state *schemaCacheTestState) invalidateLatest() { + state.mu.Lock() + state.invalid[state.latestID] = true + state.mu.Unlock() +} + +func (state *schemaCacheTestState) isValid(id int) bool { + state.mu.Lock() + defer state.mu.Unlock() + return !state.invalid[id] +} + +func (state *schemaCacheTestState) record(statement string) { + state.mu.Lock() + state.statements = append(state.statements, statement) + state.mu.Unlock() +} + +func (state *schemaCacheTestState) executedStatements() []string { + state.mu.Lock() + defer state.mu.Unlock() + return append([]string(nil), state.statements...) +} + +func (state *schemaCacheTestState) openCount() int { + state.mu.Lock() + defer state.mu.Unlock() + return state.opens +} + +type schemaCacheTestDriver struct { + state *schemaCacheTestState +} + +func (testDriver *schemaCacheTestDriver) Open(string) (driver.Conn, error) { + return testDriver.state.openConnection(), nil +} + +type schemaCacheTestConn struct { + state *schemaCacheTestState + id int +} + +func (*schemaCacheTestConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip } +func (*schemaCacheTestConn) Close() error { return nil } +func (*schemaCacheTestConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip } + +func (conn *schemaCacheTestConn) ExecContext(_ context.Context, statement string, _ []driver.NamedValue) (driver.Result, error) { + conn.state.record(statement) + return driver.RowsAffected(0), nil +} + +func (conn *schemaCacheTestConn) IsValid() bool { + return conn.state.isValid(conn.id) +} + +func openSchemaCacheTestDB(t *testing.T, state *schemaCacheTestState) *sql.DB { + t.Helper() + driverName := fmt.Sprintf("vastbase-schema-cache-%d", schemaCacheDriverSequence.Add(1)) + sql.Register(driverName, &schemaCacheTestDriver{state: state}) + db, err := sql.Open(driverName, "") + if err != nil { + t.Fatal(err) + } + db.SetMaxOpenConns(1) + db.SetMaxIdleConns(1) + t.Cleanup(func() { _ = db.Close() }) + return db +} + +func mustSchemaConn(t *testing.T, server *server, schema string) *sql.Conn { + t.Helper() + conn, err := server.schemaConn(context.Background(), schema) + if err != nil { + t.Fatalf("schemaConn(%q): %v", schema, err) + } + return conn +} diff --git a/agents/drivers/vastbase-go/driver.go b/agents/drivers/vastbase-go/driver.go new file mode 100644 index 000000000..1ab3179c8 --- /dev/null +++ b/agents/drivers/vastbase-go/driver.go @@ -0,0 +1,145 @@ +package main + +import ( + "database/sql" + "net/url" + "strings" + + _ "gitcode.com/opengauss/openGauss-connector-go-pq" +) + +const ( + agentKey = "vastbase" + agentSQLDriverName = "opengauss" + agentDefaultPort = 5432 + agentDriverName = "openGauss-connector-go-pq" + agentDriverVersion = "v1.0.8" +) + +type nativeURLParameter struct { + Key string + Value string +} + +var vastbaseDataTypes = append(append([]string{}, postgresDataTypes...), + "floatvector", "halfvector", "int8vector", "sparsevector", +) + +func agentDataTypes() []string { + return vastbaseDataTypes +} + +func detectAgentMode(_ *sql.DB, configuredMySQL bool) vastbaseMode { + if configuredMySQL { + return vastbaseMode{compatibilityMode: "mysql", mysqlCompat: true, postgresCatalog: true} + } + return vastbaseMode{compatibilityMode: "postgres", postgresCatalog: true} +} + +func agentSSLModeAttempts(sslMode string) []string { + return []string{sslMode} +} + +func agentInitialSSLMode(sslMode string) string { + return sslMode +} + +func agentSSLNotSupported(error) bool { + return false +} + +func isAgentJDBCURL(value string) bool { + normalized := strings.ToLower(strings.TrimSpace(value)) + return strings.HasPrefix(normalized, "jdbc:vastbase://") || strings.HasPrefix(normalized, "jdbc:postgresql://") +} + +func isAgentNativeURL(value string) bool { + normalized := strings.ToLower(strings.TrimSpace(value)) + return strings.HasPrefix(normalized, "postgres://") || strings.HasPrefix(normalized, "postgresql://") +} + +func normalizeAgentObjectSource(source string) string { + trimmed := strings.TrimSpace(source) + if !strings.HasPrefix(trimmed, "(") || !strings.HasSuffix(trimmed, ")") { + return source + } + inner := trimmed[1 : len(trimmed)-1] + if comma := strings.IndexByte(inner, ','); comma > 0 { + inner = strings.TrimSpace(inner[comma+1:]) + } + if len(inner) >= 2 && inner[0] == '"' && inner[len(inner)-1] == '"' { + inner = strings.ReplaceAll(inner[1:len(inner)-1], `""`, `"`) + } + return strings.TrimSpace(inner) +} + +func nativeURLParams(raw string) []nativeURLParameter { + parameters := make([]nativeURLParameter, 0) + for _, pair := range strings.FieldsFunc(raw, func(r rune) bool { return r == '&' || r == ';' }) { + key, value, ok := strings.Cut(pair, "=") + if !ok { + continue + } + key = strings.TrimSpace(key) + value = strings.TrimSpace(value) + if decoded, err := url.QueryUnescape(key); err == nil { + key = decoded + } + if decoded, err := url.QueryUnescape(value); err == nil { + value = decoded + } + if !isSafeParamKey(key) { + continue + } + normalizedKey, normalizedValue, include := nativeURLParam(key, value) + if include { + parameters = append(parameters, nativeURLParameter{Key: normalizedKey, Value: normalizedValue}) + } + } + return parameters +} + +func nativeURLParam(key, value string) (string, string, bool) { + switch strings.ToLower(strings.TrimSpace(key)) { + case "ssl": + if strings.EqualFold(value, "true") || value == "1" { + return "sslmode", "require", true + } + return "sslmode", "disable", true + case "sslmode": + if strings.EqualFold(value, "enable") { + value = "require" + } + return "sslmode", strings.ToLower(value), true + case "targetservertype": + switch strings.ToLower(value) { + case "master", "primary": + value = "primary" + case "slave", "secondary": + value = "standby" + case "preferslave", "prefersecondary", "prefer-standby": + value = "prefer-standby" + default: + value = "any" + } + return "target_session_attrs", value, true + case "connecttimeout", "logintimeout": + return "connect_timeout", value, true + case "applicationname": + return "application_name", value, true + case "currentschema": + return "search_path", value, true + case "loggerlevel": + return "loggerLevel", value, true + case "autosave", "enable_ce", "db_compatibility", "loadbalancehosts", "autobalance", + "protocolversion", "preparethreshold", "preparedstatementcachequeries", + "databasemetadatacachefields", "databasemetadatacachefieldsmib", "stringtype", + "batchmode", "fetchsize", "defaultrowfetchsize", "rewritebatchedinserts", "unknownlength", + "sockettimeout", "sockettimeoutinconnecting", "socketfactory", "socketfactoryarg", + "sslfactory", "sslfactoryarg", "sslhostnameverifier", "loggerfile", "loggerdir", + "tlcp", "sslenccert", "sslenckey", "connectionextrainfo", "nvarchartype": + return "", "", false + default: + return strings.TrimSpace(key), value, true + } +} diff --git a/agents/drivers/vastbase-go/go.mod b/agents/drivers/vastbase-go/go.mod new file mode 100644 index 000000000..fae064988 --- /dev/null +++ b/agents/drivers/vastbase-go/go.mod @@ -0,0 +1,12 @@ +module github.com/t8y2/dbx/agents/drivers/vastbase-go + +go 1.22 + +require gitcode.com/opengauss/openGauss-connector-go-pq v1.0.8 + +require ( + github.com/tjfoc/gmsm v1.4.1 // indirect + golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97 // indirect + golang.org/x/text v0.3.3 // indirect + golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1 // indirect +) diff --git a/agents/drivers/vastbase-go/go.sum b/agents/drivers/vastbase-go/go.sum new file mode 100644 index 000000000..be64eb40c --- /dev/null +++ b/agents/drivers/vastbase-go/go.sum @@ -0,0 +1,97 @@ +cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= +gitcode.com/opengauss/openGauss-connector-go-pq v1.0.8 h1:QQBQgXTOx7UP4krmxjBGTk/Sm4lh98GW3nRWkvxBBn4= +gitcode.com/opengauss/openGauss-connector-go-pq v1.0.8/go.mod h1:EIVrn+q7Ip07RUWAQxg6ELVeeY3+TjgAytV+WUIyHTs= +github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= +github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU= +github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw= +github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= +github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98= +github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c= +github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q= +github.com/golang/mock v1.1.1/go.mod h1:oTYuIxOrZwtPieC+H1uAHpcLFnEyAGVDL/k47Jfbm0A= +github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= +github.com/golang/protobuf v1.3.2/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= +github.com/golang/protobuf v1.3.3/go.mod h1:vzj43D7+SQXF/4pzW/hwtAqwc6iTitCiVSaWz5lYuqw= +github.com/golang/protobuf v1.4.0-rc.1/go.mod h1:ceaxUfeHdC40wWswd/P6IGgMaK3YpKi5j83Wpe3EHw8= +github.com/golang/protobuf v1.4.0-rc.1.0.20200221234624-67d41d38c208/go.mod h1:xKAWHe0F5eneWXFV3EuXVDTCmh+JuBKY0li0aMyXATA= +github.com/golang/protobuf v1.4.0-rc.2/go.mod h1:LlEzMj4AhA7rCAGe4KMBDvJI+AwstrUpVNzEA03Pprs= +github.com/golang/protobuf v1.4.0-rc.4.0.20200313231945-b860323f09d0/go.mod h1:WU3c8KckQ9AFe+yFwt9sWVRKCVIyN9cPHBJSNnbL67w= +github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvqG2KuDX0= +github.com/golang/protobuf v1.4.2/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI= +github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M= +github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= +github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= +github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= +github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo= +github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= +github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/prometheus/client_model v0.0.0-20190812154241-14fe0d1b01d4/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.7.0 h1:nwc3DEeHmmLAfoZucVR881uASk0Mfjw8xYJ99tb5CcY= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/tjfoc/gmsm v1.4.1 h1:aMe1GlZb+0bLjn+cKTPEvvn9oUEBlJitaZiiBwsbgho= +github.com/tjfoc/gmsm v1.4.1/go.mod h1:j4INPkHWMrhJb38G+J6W4Tw0AbuN8Thu3PbdVYhVcTE= +golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= +golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= +golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= +golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97 h1:/UOmuWzQfxxo9UtlXMwuQU8CMgg1eZXqTRwkSQJWKOI= +golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= +golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= +golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= +golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU= +golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= +golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= +golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= +golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= +golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= +golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= +golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= +golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= +golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= +golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= +golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= +golang.org/x/text v0.3.3 h1:cokOdA+Jmi5PJGXLlLllQSgYigAEfHXJAERHVMaCc2k= +golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= +golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= +golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY= +golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= +golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q= +golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= +golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1 h1:go1bK/D/BFZV2I8cIQd1NKEZ+0owSTG1fDTci4IqFcE= +golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= +google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM= +google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4= +google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc= +google.golang.org/genproto v0.0.0-20190819201941-24fa4b261c55/go.mod h1:DMBHOl98Agz4BDEuKkezgsaosCRResVns1a3J2ZsMNc= +google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c= +google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyacEbxg= +google.golang.org/grpc v1.25.1/go.mod h1:c3i+UQWmh7LiEpx4sFZnkU36qjEYZ0imhYfXVyQciAY= +google.golang.org/grpc v1.31.0/go.mod h1:N36X2cJ7JwdamYAgDz+s+rVMFjt3numwzf/HckM8pak= +google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= +google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= +google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM= +google.golang.org/protobuf v1.20.1-0.20200309200217-e05f789c0967/go.mod h1:A+miEFZTKqfCUM6K7xSMQL9OKL/b6hQv+e19PK+JZNE= +google.golang.org/protobuf v1.21.0/go.mod h1:47Nbq4nVaFHyn7ilMalzfO3qCViNmqZ2kzikPIcrTAo= +google.golang.org/protobuf v1.23.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= +honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= diff --git a/agents/drivers/vastbase-go/integration_test.go b/agents/drivers/vastbase-go/integration_test.go new file mode 100644 index 000000000..6cd5780e0 --- /dev/null +++ b/agents/drivers/vastbase-go/integration_test.go @@ -0,0 +1,156 @@ +package main + +import ( + "encoding/json" + "fmt" + "os" + "strconv" + "strings" + "testing" + "time" +) + +func TestVastbaseIntegration(t *testing.T) { + host := os.Getenv("VASTBASE_TEST_HOST") + portText := os.Getenv("VASTBASE_TEST_PORT") + username := os.Getenv("VASTBASE_TEST_USERNAME") + password := os.Getenv("VASTBASE_TEST_PASSWORD") + if host == "" || portText == "" || username == "" || password == "" { + t.Skip("Vastbase integration environment is not configured") + } + port, err := strconv.Atoi(portText) + if err != nil { + t.Fatal(err) + } + database := os.Getenv("VASTBASE_TEST_DATABASE") + if database == "" { + database = "test" + } + suffix := strconv.FormatInt(time.Now().UnixNano(), 36) + parent := "dbx_go_parent_" + suffix + child := "dbx_go_child_" + suffix + view := "dbx_go_view_" + suffix + function := "dbx_go_fn_" + suffix + + server := newServer() + cp := connectParams{ + Host: host, Port: port, Database: database, Username: username, Password: password, + ConnectionString: fmt.Sprintf("jdbc:vastbase://%s:%d/%s", host, port, database), + } + if err := server.connect(cp); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = server.disconnect() }) + cleanup := []string{ + "DROP VIEW IF EXISTS public." + quoteIdentifier(view), + "DROP FUNCTION IF EXISTS public." + quoteIdentifier(function) + "()", + "DROP TABLE IF EXISTS public." + quoteIdentifier(child), + "DROP TABLE IF EXISTS public." + quoteIdentifier(parent), + } + t.Cleanup(func() { + for _, statement := range cleanup { + _, _ = server.executeQuery(queryOptions{SQL: statement}) + } + }) + + mustExecute(t, server, "CREATE TABLE public."+quoteIdentifier(parent)+" (id integer PRIMARY KEY, name varchar(64) NOT NULL)") + mustExecute(t, server, "COMMENT ON TABLE public."+quoteIdentifier(parent)+" IS '订单父表'") + mustExecute(t, server, "COMMENT ON COLUMN public."+quoteIdentifier(parent)+".id IS '主键编号'") + mustExecute(t, server, "COMMENT ON COLUMN public."+quoteIdentifier(parent)+".name IS '客户''名称'") + mustExecute(t, server, "CREATE TABLE public."+quoteIdentifier(child)+" (id integer PRIMARY KEY, parent_id integer REFERENCES public."+quoteIdentifier(parent)+"(id))") + mustExecute(t, server, "CREATE INDEX "+quoteIdentifier(child+"_parent_idx")+" ON public."+quoteIdentifier(child)+"(parent_id)") + mustExecute(t, server, "CREATE VIEW public."+quoteIdentifier(view)+" AS SELECT id, name FROM public."+quoteIdentifier(parent)) + mustExecute(t, server, "CREATE FUNCTION public."+quoteIdentifier(function)+"() RETURNS text AS $$ SELECT 'dbx'; $$ LANGUAGE SQL") + + tables, err := server.listTables("public", metadataListConstraints{Filter: suffix}) + if err != nil || len(tables) < 3 { + t.Fatalf("list tables failed: count=%d err=%v", len(tables), err) + } + columns, err := server.getColumns("public", child) + if err != nil || len(columns) != 2 || !columns[0].IsPrimaryKey { + t.Fatalf("get columns failed: columns=%v err=%v", columns, err) + } + parentColumns, err := server.getColumns("public", parent) + if err != nil || len(parentColumns) != 2 || parentColumns[0].Comment == nil || *parentColumns[0].Comment != "主键编号" || parentColumns[1].Comment == nil || *parentColumns[1].Comment != "客户'名称" { + t.Fatalf("get commented columns failed: columns=%v err=%v", parentColumns, err) + } + ddl, err := server.getTableDDL("public", parent) + if err != nil { + t.Fatalf("get table DDL failed: %v", err) + } + qualifiedParent := quoteIdentifier("public") + "." + quoteIdentifier(parent) + for _, expected := range []string{ + "COMMENT ON TABLE " + qualifiedParent + " IS '订单父表';", + "COMMENT ON COLUMN " + qualifiedParent + "." + quoteIdentifier("id") + " IS '主键编号';", + "COMMENT ON COLUMN " + qualifiedParent + "." + quoteIdentifier("name") + " IS '客户''名称';", + } { + if !strings.Contains(ddl, expected) { + t.Fatalf("table DDL missing %q:\n%s", expected, ddl) + } + } + indexes, err := server.listIndexes("public", child) + if err != nil || len(indexes) < 2 { + t.Fatalf("list indexes failed: indexes=%v err=%v", indexes, err) + } + foreignKeys, err := server.listForeignKeys("public", child) + if err != nil || len(foreignKeys) != 1 || foreignKeys[0].RefTable != parent { + t.Fatalf("list foreign keys failed: keys=%v err=%v", foreignKeys, err) + } + source, err := server.getObjectSource("public", function, "FUNCTION") + if err != nil || !strings.Contains(fmt.Sprint(source["source"]), function) { + t.Fatalf("get function source failed: source=%v err=%v", source, err) + } + + transactionParams := map[string]json.RawMessage{ + "schema": rawJSON("public"), + "statements": rawJSON([]string{"INSERT INTO " + quoteIdentifier(parent) + " VALUES (1, 'one')", "INSERT INTO " + quoteIdentifier(child) + " VALUES (1, 1)"}), + } + if _, err := server.executeTransaction(transactionParams); err != nil { + t.Fatal(err) + } + page, err := server.executeQueryPage(queryOptions{SQL: "SELECT generate_series(1, 250)", MaxRows: 250}, 100) + if err != nil || !page.HasMore || page.SessionID == nil || len(page.Rows) != 100 { + t.Fatalf("first page failed: page=%v err=%v", page, err) + } + second, err := server.fetchQueryPage(*page.SessionID, 100) + if err != nil || !second.HasMore || len(second.Rows) != 100 { + t.Fatalf("second page failed: page=%v err=%v", second, err) + } + third, err := server.fetchQueryPage(*page.SessionID, 100) + if err != nil || third.HasMore || len(third.Rows) != 50 { + t.Fatalf("third page failed: page=%v err=%v", third, err) + } + + cancelStart := time.Now() + cancelResult := make(chan error, 1) + go func() { + _, queryErr := server.executeQuery(queryOptions{SQL: "SELECT pg_sleep(5)", MaxRows: 1}) + cancelResult <- queryErr + }() + time.Sleep(200 * time.Millisecond) + server.cancelActiveQuery() + if queryErr := <-cancelResult; queryErr == nil { + t.Fatal("cancel_session did not interrupt the active query") + } + if elapsed := time.Since(cancelStart); elapsed > 3*time.Second { + t.Fatalf("query cancellation was too slow: %s", elapsed) + } + if err := server.validateConnection(); err != nil { + t.Fatalf("connection was not reusable after cancellation: %v", err) + } +} + +func rawJSON(value any) json.RawMessage { + data, err := json.Marshal(value) + if err != nil { + panic(err) + } + return json.RawMessage(data) +} + +func mustExecute(t *testing.T, server *server, statement string) { + t.Helper() + if _, err := server.executeQuery(queryOptions{SQL: statement}); err != nil { + t.Fatalf("execute %q: %v", statement, err) + } +} diff --git a/agents/drivers/vastbase-go/main.go b/agents/drivers/vastbase-go/main.go new file mode 100644 index 000000000..7885bb044 --- /dev/null +++ b/agents/drivers/vastbase-go/main.go @@ -0,0 +1,1349 @@ +package main + +import ( + "bufio" + "context" + "database/sql" + "database/sql/driver" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "net/url" + "os" + "strings" + "sync" + "time" +) + +const ( + protocolVersion = 2 + defaultMaxRows = 10000 + legacyAgentSessionID = "__legacy__" + maxAgentSessions = 256 + defaultConnectTimeout = 15 * time.Second +) + +type request struct { + ID json.RawMessage `json:"id"` + Method string `json:"method"` + Params map[string]json.RawMessage `json:"params"` +} + +type response struct { + JSONRPC string `json:"jsonrpc,omitempty"` + ID json.RawMessage `json:"id,omitempty"` + Result any `json:"result,omitempty"` + Error *rpcError `json:"error,omitempty"` +} + +type connectParams struct { + Host string `json:"host"` + Port int `json:"port"` + Database string `json:"database"` + Username string `json:"username"` + Password string `json:"password"` + URLParams string `json:"url_params"` + ConnectionString string `json:"connection_string"` + MySQLCompatMode bool `json:"mysql_compat_mode"` + SSL bool `json:"ssl"` + CACertPath string `json:"ca_cert_path"` + ClientCertPath string `json:"client_cert_path"` + ClientKeyPath string `json:"client_key_path"` + SessionRole string `json:"sessionRole"` +} + +type queryOptions struct { + SQL string `json:"sql"` + Database string `json:"database"` + Schema string `json:"schema"` + MaxRows int `json:"maxRows"` + FetchSize int `json:"fetchSize"` + TimeoutSecs int `json:"timeoutSecs"` +} + +type completionAssistantRequest struct { + ConnectionID string `json:"connection_id"` + Database string `json:"database"` + Schema string `json:"schema"` + ObjectKinds []string `json:"object_kinds"` + Mask string `json:"mask"` + CaseSensitive bool `json:"case_sensitive"` + GlobalSearch bool `json:"global_search"` + MaxResults int `json:"max_results"` + ParentSchema string `json:"parent_schema"` + ParentName string `json:"parent_name"` + MatchMode string `json:"match_mode"` +} + +type completionAssistantCandidate struct { + Name string `json:"name"` + Kind string `json:"kind"` + Database *string `json:"database"` + Schema *string `json:"schema"` + ParentSchema *string `json:"parent_schema"` + ParentName *string `json:"parent_name"` + Comment *string `json:"comment"` + DataType *string `json:"data_type"` +} + +type completionAssistantResponse struct { + Candidates []completionAssistantCandidate `json:"candidates"` + Incomplete bool `json:"incomplete"` + FallbackUsed bool `json:"fallback_used"` +} + +type queryResult struct { + Columns []string `json:"columns"` + ColumnTypes []string `json:"column_types"` + SpatialColumns []spatialColumn `json:"spatial_columns,omitempty"` + SpatialValues [][]*uint32 `json:"spatial_values,omitempty"` + Rows [][]any `json:"rows"` + AffectedRows int64 `json:"affected_rows"` + ExecutionTimeMS int64 `json:"execution_time_ms"` + Truncated bool `json:"truncated"` +} + +type queryPageResult struct { + Columns []string `json:"columns"` + ColumnTypes []string `json:"column_types"` + SpatialColumns []spatialColumn `json:"spatial_columns,omitempty"` + SpatialValues [][]*uint32 `json:"spatial_values,omitempty"` + Rows [][]any `json:"rows"` + AffectedRows int64 `json:"affected_rows"` + ExecutionTimeMS int64 `json:"execution_time_ms"` + Truncated bool `json:"truncated"` + SessionID *string `json:"session_id"` + HasMore bool `json:"has_more"` +} + +type querySession struct { + rows *sql.Rows + conn *sql.Conn + columns []string + columnTypes []string + scanner *rowScanner + pending []any + pendingSpatial []*uint32 + remaining int + cancel context.CancelFunc +} + +type rowScanner struct { + values []any + destinations []any + spatial *spatialDecoder +} + +type server struct { + db *sql.DB + openDatabase agentDBOpener + params connectParams + mode vastbaseMode + usePgDefaultExpression bool + catalogIdentityUnsupported bool + infoColumnTypeUnsupported bool + infoUdtNameUnsupported bool + listTablesStatement *sql.Stmt + connectionRuntime *connectionRuntime + sessionAffinity bool + currentSchema string + schemaInitialized bool + schemaConnectionID uintptr + sessions map[string]*querySession + nextSessionID uint64 + activeCancelMu sync.Mutex + activeCancel context.CancelFunc +} + +type agentSession struct { + server *server + runtimeKey string + mu sync.Mutex +} + +type runtimeServer struct { + mu sync.RWMutex + sessions map[string]*agentSession + connectionRuntimeMu sync.Mutex + connectionRuntimes map[string]*connectionRuntime +} + +func main() { + runtime := &runtimeServer{sessions: map[string]*agentSession{}} + encoder := json.NewEncoder(os.Stdout) + var encoderMu sync.Mutex + var requests sync.WaitGroup + fmt.Fprintln(os.Stdout, `{"ready":true}`) + + scanner := bufio.NewScanner(os.Stdin) + scanner.Buffer(make([]byte, 0, 64*1024), 512*1024*1024) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + var envelope request + if json.Unmarshal([]byte(line), &envelope) == nil && envelope.Method == "shutdown" { + requests.Wait() + resp, _ := runtime.handleLine(line) + encoderMu.Lock() + _ = encoder.Encode(resp) + encoderMu.Unlock() + return + } + requests.Add(1) + go func(line string) { + defer requests.Done() + resp, _ := runtime.handleLine(line) + encoderMu.Lock() + defer encoderMu.Unlock() + if err := encoder.Encode(resp); err != nil { + fmt.Fprintf(os.Stderr, "failed to write response: %v\n", err) + } + }(line) + } + requests.Wait() +} + +func (r *runtimeServer) handleLine(line string) (response, bool) { + var req request + if err := json.Unmarshal([]byte(line), &req); err != nil { + return errorResponse(nil, "", "", err), false + } + if len(req.ID) == 0 { + req.ID = json.RawMessage("1") + } + result, shutdown, err := r.dispatch(req.Method, req.Params) + if err != nil { + return errorResponse(req.ID, req.Method, stringParam(req.Params, "agentSessionId"), err), false + } + return response{JSONRPC: "2.0", ID: req.ID, Result: result}, shutdown +} + +func (r *runtimeServer) dispatch(method string, params map[string]json.RawMessage) (any, bool, error) { + switch method { + case "handshake": + return map[string]any{ + "protocolVersion": protocolVersion, + "agentProtocolVersion": protocolVersion, + "capabilities": []string{ + "connect", "test_connection", "metadata", "query", "paged_query", "transaction", "ddl", "multi_session", "structured_error_v1", + }, + }, false, nil + case "open_session": + id := stringParam(params, "agentSessionId") + if id == "" { + return nil, false, errors.New("agentSessionId is required") + } + var cp connectParams + if err := decodeParams(params, &cp); err != nil { + return nil, false, err + } + return map[string]bool{"ok": true}, false, r.openSession(id, cp) + case "close_session": + return map[string]bool{"ok": true}, false, r.closeSession(stringParam(params, "agentSessionId")) + case "validate_session": + session, err := r.session(stringParam(params, "agentSessionId")) + if err != nil { + return nil, false, err + } + session.mu.Lock() + defer session.mu.Unlock() + release, permitErr := session.server.acquireOperationPermit("validate_session") + if permitErr != nil { + return nil, false, permitErr + } + defer release() + return map[string]bool{"ok": true}, false, session.server.validateConnection() + case "cancel_session": + session, err := r.session(stringParam(params, "agentSessionId")) + if err != nil { + return nil, false, err + } + session.server.cancelActiveQuery() + return map[string]bool{"ok": true}, false, nil + case "test_connection": + return newServer().dispatch(method, params) + case "connect": + var cp connectParams + if err := decodeParams(params, &cp); err != nil { + return nil, false, err + } + _ = r.closeSession(legacyAgentSessionID) + return map[string]bool{"ok": true}, false, r.openSession(legacyAgentSessionID, cp) + case "disconnect": + return map[string]bool{"ok": true}, false, r.closeSession(legacyAgentSessionID) + case "shutdown": + return map[string]bool{"ok": true}, true, r.closeAllSessions() + default: + id := stringParam(params, "agentSessionId") + if id == "" { + id = legacyAgentSessionID + } + session, err := r.session(id) + if err != nil { + return nil, false, err + } + session.mu.Lock() + defer session.mu.Unlock() + release, permitErr := session.server.acquireOperationPermit(method) + if permitErr != nil { + return nil, false, permitErr + } + defer release() + return session.server.dispatch(method, params) + } +} + +func (r *runtimeServer) openSession(id string, cp connectParams) error { + r.mu.Lock() + if _, exists := r.sessions[id]; exists { + r.mu.Unlock() + return fmt.Errorf("agent session already exists: %s", id) + } + if len(r.sessions) >= maxAgentSessions { + r.mu.Unlock() + return fmt.Errorf("agent session limit reached: %d", maxAgentSessions) + } + r.mu.Unlock() + + connectionRuntime, runtimeKey := r.acquireConnectionRuntime(cp) + s := newServer() + if err := s.connectWithRuntime(cp, connectionRuntime); err != nil { + r.releaseConnectionRuntime(runtimeKey) + return err + } + r.mu.Lock() + defer r.mu.Unlock() + if _, exists := r.sessions[id]; exists { + _ = s.disconnect() + r.releaseConnectionRuntime(runtimeKey) + return fmt.Errorf("agent session already exists: %s", id) + } + r.sessions[id] = &agentSession{server: s, runtimeKey: runtimeKey} + return nil +} + +func (r *runtimeServer) session(id string) (*agentSession, error) { + r.mu.RLock() + session := r.sessions[id] + r.mu.RUnlock() + if session == nil { + return nil, fmt.Errorf("agent session not found: %s", id) + } + return session, nil +} + +func (r *runtimeServer) closeSession(id string) error { + r.mu.Lock() + session := r.sessions[id] + delete(r.sessions, id) + r.mu.Unlock() + if session == nil { + return nil + } + session.server.cancelActiveQuery() + session.mu.Lock() + defer session.mu.Unlock() + err := session.server.disconnect() + r.releaseConnectionRuntime(session.runtimeKey) + return err +} + +func (r *runtimeServer) closeAllSessions() error { + r.mu.RLock() + ids := make([]string, 0, len(r.sessions)) + for id := range r.sessions { + ids = append(ids, id) + } + r.mu.RUnlock() + var firstErr error + for _, id := range ids { + if err := r.closeSession(id); err != nil && firstErr == nil { + firstErr = err + } + } + if err := r.closeConnectionRuntimes(); err != nil && firstErr == nil { + firstErr = err + } + return firstErr +} + +func newServer() *server { + return &server{openDatabase: openDBWithSSLMode, sessions: map[string]*querySession{}} +} + +func (s *server) dispatch(method string, params map[string]json.RawMessage) (any, bool, error) { + switch method { + case "handshake": + return map[string]any{ + "protocolVersion": protocolVersion, + "agentProtocolVersion": protocolVersion, + "capabilities": []string{"connect", "test_connection", "metadata", "query", "paged_query", "transaction", "ddl", "structured_error_v1"}, + }, false, nil + case "connect": + var cp connectParams + if err := decodeParams(params, &cp); err != nil { + return nil, false, err + } + return map[string]bool{"ok": true}, false, s.connect(cp) + case "test_connection": + var cp connectParams + if err := decodeParams(params, &cp); err != nil { + return nil, false, err + } + result, err := s.testConnection(cp) + return result, false, err + case "validate_connection": + return map[string]bool{"ok": true}, false, s.validateConnection() + case "connection_info": + info, err := s.connectionInfo() + return info, false, err + case "list_databases": + result, err := s.listDatabases() + return result, false, err + case "list_schemas": + result, err := s.listSchemas(stringSliceParam(params, "visible_schemas"), boolParam(params, "show_system_schemas")) + return result, false, err + case "list_tables": + result, err := s.listTables(stringParam(params, "schema"), metadataListConstraintsFromParams(params)) + return result, false, err + case "get_table_comment": + result, err := s.getTableComment(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "list_objects": + result, err := s.listObjects(stringParam(params, "schema"), metadataListConstraintsFromParams(params)) + return result, false, err + case "list_data_types": + return agentDataTypes(), false, nil + case "completion_assistant_search_v1": + var request completionAssistantRequest + if err := decodeParams(params, &request); err != nil { + return nil, false, err + } + result, err := s.completionAssistantSearch(request) + return result, false, err + case "get_columns": + result, err := s.getColumns(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "list_indexes": + result, err := s.listIndexes(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "list_foreign_keys": + result, err := s.listForeignKeys(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "list_triggers": + result, err := s.listTriggers(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "get_object_source": + result, err := s.getObjectSource(stringParam(params, "schema"), stringParam(params, "name"), stringParam(params, "object_type")) + return result, false, err + case "get_table_ddl": + result, err := s.getTableDDL(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "get_explain_info": + result, err := s.getExplainInfo(stringParam(params, "sql")) + return map[string]any{"plan": result, "has_actual_stats": false}, false, err + case "execute_query": + opts := queryOptionsFromParams(params) + result, err := s.executeQuery(opts) + return result, false, err + case "execute_query_page", "start_table_read": + opts := queryOptionsFromParams(params) + result, err := s.executeQueryPage(opts, intParam(params, "pageSize")) + return result, false, err + case "fetch_query_page", "fetch_table_read_page": + result, err := s.fetchQueryPage(stringParam(params, "sessionId"), intParam(params, "pageSize")) + return result, false, err + case "close_query_session", "close_table_read_session": + return s.closeQuerySession(stringParam(params, "sessionId")), false, nil + case "execute_transaction": + result, err := s.executeTransaction(params) + return result, false, err + case "execute_batch": + result, err := s.executeBatch(params) + return result, false, err + case "disconnect": + return map[string]bool{"ok": true}, false, s.disconnect() + case "shutdown": + return map[string]bool{"ok": true}, true, s.disconnect() + default: + return nil, false, fmt.Errorf("unknown method: %s", method) + } +} + +func (s *server) connect(cp connectParams) error { + _ = s.disconnect() + db, err := openAndPingDB(cp, defaultConnectTimeout, s.openDatabase) + if err != nil { + return err + } + s.db = db + s.params = cp + s.mode = detectAgentMode(db, cp.MySQLCompatMode) + s.usePgDefaultExpression = false + s.catalogIdentityUnsupported = false + s.infoColumnTypeUnsupported = false + s.infoUdtNameUnsupported = false + s.sessionAffinity = false + return nil +} + +func (s *server) connectWithRuntime(cp connectParams, connectionRuntime *connectionRuntime) error { + _ = s.disconnect() + if err := connectionRuntime.validate(cp, s.openDatabase); err != nil { + return err + } + db, err := s.openDatabase(cp, agentInitialSSLMode(effectiveSSLMode(cp))) + if err != nil { + return err + } + s.db = db + s.connectionRuntime = connectionRuntime + s.params = cp + s.mode = detectAgentMode(db, cp.MySQLCompatMode) + s.usePgDefaultExpression = false + s.catalogIdentityUnsupported = false + s.infoColumnTypeUnsupported = false + s.infoUdtNameUnsupported = false + s.sessionAffinity = false + return nil +} + +func (s *server) testConnection(cp connectParams) (map[string]any, error) { + db, err := openAndPingDB(cp, defaultConnectTimeout, s.openDatabase) + if err != nil { + return nil, err + } + defer db.Close() + temporary := newServer() + temporary.db = db + temporary.params = cp + temporary.mode = detectAgentMode(db, cp.MySQLCompatMode) + info, err := temporary.connectionInfo() + if err != nil { + return nil, err + } + result := map[string]any{"ok": true} + if databaseInfo, ok := info["databaseInfo"]; ok { + result["databaseInfo"] = databaseInfo + } + return result, nil +} + +type agentDBOpener func(connectParams, string) (*sql.DB, error) + +func openAndPingDB(cp connectParams, timeout time.Duration, opener agentDBOpener) (*sql.DB, error) { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + sslMode := effectiveSSLMode(cp) + attempts := agentSSLModeAttempts(sslMode) + for index, attempt := range attempts { + db, err := opener(cp, attempt) + if err == nil { + err = db.PingContext(ctx) + } + if err == nil { + return db, nil + } + if db != nil { + _ = db.Close() + } + if index == 0 && len(attempts) > 1 && agentSSLNotSupported(err) { + continue + } + return nil, err + } + return nil, fmt.Errorf("%s connection failed", agentKey) +} + +func openDBWithSSLMode(cp connectParams, sslMode string) (*sql.DB, error) { + dsn := buildDSNWithSSLMode(cp, sslMode) + db, err := sql.Open(agentSQLDriverName, dsn) + if err != nil { + return nil, err + } + // Each protocol session is serialized and owns one database connection. + // Keeping a single physical connection preserves session state such as + // search_path and avoids extra pool coordination on the hot query path. + db.SetMaxOpenConns(1) + db.SetMaxIdleConns(1) + db.SetConnMaxLifetime(5 * time.Minute) + return db, nil +} + +func (s *server) disconnect() error { + s.cancelActiveQuery() + s.closeAllQuerySessions() + s.usePgDefaultExpression = false + s.catalogIdentityUnsupported = false + s.infoColumnTypeUnsupported = false + s.infoUdtNameUnsupported = false + s.connectionRuntime = nil + s.sessionAffinity = false + s.resetSchemaCache() + if s.listTablesStatement != nil { + _ = s.listTablesStatement.Close() + s.listTablesStatement = nil + } + if s.db == nil { + return nil + } + err := s.db.Close() + s.db = nil + return err +} + +func (s *server) validateConnection() error { + db, err := s.metadataDatabase() + if err != nil { + return err + } + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + for attempt := 0; attempt < 2; attempt++ { + err = db.PingContext(ctx) + if !errors.Is(err, driver.ErrBadConn) { + return err + } + } + return err +} + +func (s *server) requireDB() (*sql.DB, error) { + if s.db == nil { + return nil, errors.New("not connected") + } + return s.db, nil +} + +func (s *server) beginOperation(timeoutSecs int) (context.Context, context.CancelFunc) { + ctx := context.Background() + var cancel context.CancelFunc + if timeoutSecs > 0 { + ctx, cancel = context.WithTimeout(ctx, time.Duration(timeoutSecs)*time.Second) + } else { + ctx, cancel = context.WithCancel(ctx) + } + s.activeCancelMu.Lock() + s.activeCancel = cancel + s.activeCancelMu.Unlock() + return ctx, cancel +} + +func (s *server) endOperation(cancel context.CancelFunc) { + cancel() + s.activeCancelMu.Lock() + s.activeCancel = nil + s.activeCancelMu.Unlock() +} + +func (s *server) cancelActiveQuery() { + s.activeCancelMu.Lock() + cancel := s.activeCancel + s.activeCancelMu.Unlock() + if cancel != nil { + cancel() + } +} + +func (s *server) executeQuery(opts queryOptions) (queryResult, error) { + start := time.Now() + sqlText := trimStatementSQL(opts.SQL) + defer s.noteSQLSessionState(sqlText) + if isQuerySQL(sqlText) { + rows, conn, cancel, err := s.queryRows(sqlText, opts.Schema, opts.TimeoutSecs) + if err != nil { + return queryResult{}, err + } + defer func() { + _ = rows.Close() + _ = conn.Close() + s.endOperation(cancel) + }() + maxRows := opts.MaxRows + if maxRows <= 0 { + maxRows = defaultMaxRows + } + result, err := readRows(rows, maxRows) + result.ExecutionTimeMS = time.Since(start).Milliseconds() + return result, err + } + conn, ctx, cancel, err := s.operationConn(opts.Schema, opts.TimeoutSecs) + if err != nil { + return queryResult{}, err + } + defer func() { + _ = conn.Close() + s.endOperation(cancel) + }() + execResult, err := conn.ExecContext(ctx, sqlText) + if err != nil { + return queryResult{}, err + } + affected, _ := execResult.RowsAffected() + return queryResult{Columns: []string{}, ColumnTypes: []string{}, Rows: [][]any{}, AffectedRows: affected, ExecutionTimeMS: time.Since(start).Milliseconds()}, nil +} + +func (s *server) queryRows(sqlText string, schema string, timeoutSecs int) (*sql.Rows, *sql.Conn, context.CancelFunc, error) { + conn, ctx, cancel, err := s.operationConn(schema, timeoutSecs) + if err != nil { + return nil, nil, nil, err + } + rows, err := conn.QueryContext(ctx, sqlText) + if err != nil { + _ = conn.Close() + s.endOperation(cancel) + return nil, nil, nil, err + } + return rows, conn, cancel, nil +} + +func (s *server) executeQueryPage(opts queryOptions, pageSize int) (queryPageResult, error) { + start := time.Now() + sqlText := trimStatementSQL(opts.SQL) + defer s.noteSQLSessionState(sqlText) + if !isQuerySQL(sqlText) { + result, err := s.executeQuery(opts) + return queryPageResult{Columns: result.Columns, ColumnTypes: result.ColumnTypes, SpatialColumns: result.SpatialColumns, SpatialValues: result.SpatialValues, Rows: result.Rows, AffectedRows: result.AffectedRows, ExecutionTimeMS: result.ExecutionTimeMS, Truncated: result.Truncated}, err + } + rows, conn, cancel, err := s.queryRows(sqlText, opts.Schema, opts.TimeoutSecs) + if err != nil { + return queryPageResult{}, err + } + columns, err := rows.Columns() + if err != nil { + _ = rows.Close() + _ = conn.Close() + s.endOperation(cancel) + return queryPageResult{}, err + } + maxRows := opts.MaxRows + if maxRows <= 0 { + maxRows = defaultMaxRows + } + columnTypes := columnTypeNames(rows) + session := &querySession{rows: rows, conn: conn, columns: columns, columnTypes: columnTypes, scanner: newRowScanner(len(columns), newSpatialDecoder(columnTypes)), remaining: maxRows, cancel: cancel} + result, err := readQuerySessionPage(session, pageSize) + result.ExecutionTimeMS = time.Since(start).Milliseconds() + if err != nil { + _ = rows.Close() + _ = conn.Close() + s.endOperation(cancel) + return queryPageResult{}, err + } + if result.HasMore { + s.nextSessionID++ + id := fmt.Sprintf("%s-%d", agentKey, s.nextSessionID) + s.sessions[id] = session + result.SessionID = &id + } else { + _ = rows.Close() + _ = conn.Close() + s.endOperation(cancel) + } + return result, nil +} + +func (s *server) fetchQueryPage(id string, pageSize int) (queryPageResult, error) { + session := s.sessions[id] + if session == nil { + return queryPageResult{Columns: []string{}, ColumnTypes: []string{}, Rows: [][]any{}}, nil + } + result, err := readQuerySessionPage(session, pageSize) + if err != nil { + s.closeQuerySession(id) + return queryPageResult{}, err + } + if result.HasMore { + result.SessionID = &id + } else { + s.closeQuerySession(id) + } + return result, nil +} + +func (s *server) closeQuerySession(id string) bool { + session := s.sessions[id] + if session == nil { + return false + } + _ = session.rows.Close() + if session.conn != nil { + _ = session.conn.Close() + } + if session.cancel != nil { + s.endOperation(session.cancel) + } + delete(s.sessions, id) + return true +} + +func (s *server) closeAllQuerySessions() { + for id := range s.sessions { + s.closeQuerySession(id) + } +} + +func readQuerySessionPage(session *querySession, pageSize int) (queryPageResult, error) { + if pageSize <= 0 { + pageSize = 100 + } + capacity := min(pageSize, session.remaining) + result := queryPageResult{Columns: session.columns, ColumnTypes: session.columnTypes, Rows: make([][]any, 0, capacity)} + spatialValues := make([][]*uint32, 0, capacity) + for len(result.Rows) < pageSize && session.remaining > 0 { + if session.pending != nil { + result.Rows = append(result.Rows, session.pending) + if session.scanner.spatial != nil { + spatialValues = append(spatialValues, session.pendingSpatial) + } + session.pending = nil + session.pendingSpatial = nil + session.remaining-- + continue + } + if !session.rows.Next() { + return finishSpatialPage(result, session.scanner.spatial, spatialValues), session.rows.Err() + } + row, rowSpatial, err := session.scanner.scan(session.rows) + if err != nil { + return queryPageResult{}, err + } + result.Rows = append(result.Rows, row) + if session.scanner.spatial != nil { + spatialValues = append(spatialValues, rowSpatial) + } + session.remaining-- + } + if session.remaining <= 0 { + result.Truncated = true + return finishSpatialPage(result, session.scanner.spatial, spatialValues), nil + } + if session.rows.Next() { + row, rowSpatial, err := session.scanner.scan(session.rows) + if err != nil { + return queryPageResult{}, err + } + session.pending = row + session.pendingSpatial = rowSpatial + result.HasMore = true + } + return finishSpatialPage(result, session.scanner.spatial, spatialValues), session.rows.Err() +} + +func readRows(rows *sql.Rows, maxRows int) (queryResult, error) { + columns, err := rows.Columns() + if err != nil { + return queryResult{}, err + } + columnTypes := columnTypeNames(rows) + spatial := newSpatialDecoder(columnTypes) + scanner := newRowScanner(len(columns), spatial) + result := queryResult{Columns: columns, ColumnTypes: columnTypes, Rows: make([][]any, 0, min(maxRows, 1024))} + spatialValues := make([][]*uint32, 0, min(maxRows, 1024)) + for rows.Next() { + if len(result.Rows) >= maxRows { + result.Truncated = true + break + } + row, rowSpatial, err := scanner.scan(rows) + if err != nil { + return queryResult{}, err + } + result.Rows = append(result.Rows, row) + if spatial != nil { + spatialValues = append(spatialValues, rowSpatial) + } + } + result.SpatialColumns, result.SpatialValues = spatialResultMetadata(spatial, spatialValues) + return result, rows.Err() +} + +func newRowScanner(count int, spatial *spatialDecoder) *rowScanner { + scanner := &rowScanner{ + values: make([]any, count), + destinations: make([]any, count), + spatial: spatial, + } + for index := range scanner.values { + scanner.destinations[index] = &scanner.values[index] + } + return scanner +} + +func (scanner *rowScanner) scan(rows *sql.Rows) ([]any, []*uint32, error) { + if err := rows.Scan(scanner.destinations...); err != nil { + return nil, nil, err + } + result := make([]any, len(scanner.values)) + copy(result, scanner.values) + if scanner.spatial != nil { + return scanner.spatial.normalizeRow(result) + } + for index, value := range result { + result[index] = normalizeValue(value) + } + return result, nil, nil +} + +func columnTypeNames(rows *sql.Rows) []string { + types, err := rows.ColumnTypes() + if err != nil { + return []string{} + } + result := make([]string, len(types)) + for i, columnType := range types { + result[i] = columnType.DatabaseTypeName() + } + return result +} + +func (s *server) executeTransaction(params map[string]json.RawMessage) (queryResult, error) { + statements := stringSliceParam(params, "statements") + defer func() { + for _, statement := range statements { + s.noteSQLSessionState(statement) + } + }() + conn, ctx, cancel, err := s.operationConn(stringParam(params, "schema"), intParam(params, "timeoutSecs")) + if err != nil { + return queryResult{}, err + } + defer func() { + _ = conn.Close() + s.endOperation(cancel) + }() + start := time.Now() + tx, err := conn.BeginTx(ctx, nil) + if err != nil { + return queryResult{}, err + } + var affected int64 + for _, statement := range statements { + result, execErr := tx.ExecContext(ctx, trimStatementSQL(statement)) + if execErr != nil { + _ = tx.Rollback() + return queryResult{}, execErr + } + rows, _ := result.RowsAffected() + affected += rows + } + if err := tx.Commit(); err != nil { + return queryResult{}, err + } + return queryResult{Columns: []string{}, ColumnTypes: []string{}, Rows: [][]any{}, AffectedRows: affected, ExecutionTimeMS: time.Since(start).Milliseconds()}, nil +} + +func (s *server) executeBatch(params map[string]json.RawMessage) (queryResult, error) { + start := time.Now() + var affected int64 + for _, statement := range stringSliceParam(params, "statements") { + result, err := s.executeQuery(queryOptions{SQL: statement, Schema: stringParam(params, "schema")}) + if err != nil { + return queryResult{}, err + } + affected += result.AffectedRows + } + return queryResult{Columns: []string{}, ColumnTypes: []string{}, Rows: [][]any{}, AffectedRows: affected, ExecutionTimeMS: time.Since(start).Milliseconds()}, nil +} + +func (s *server) operationConn(schema string, timeoutSecs int) (*sql.Conn, context.Context, context.CancelFunc, error) { + ctx, cancel := s.beginOperation(timeoutSecs) + conn, err := s.schemaConn(ctx, schema) + if err != nil { + s.endOperation(cancel) + return nil, nil, nil, err + } + return conn, ctx, cancel, nil +} + +func (s *server) schemaConn(ctx context.Context, schema string) (*sql.Conn, error) { + db, err := s.requireDB() + if err != nil { + return nil, err + } + for attempt := 0; attempt < 2; attempt++ { + conn, connErr := db.Conn(ctx) + if connErr != nil { + if attempt == 0 && errors.Is(connErr, driver.ErrBadConn) { + continue + } + return nil, connErr + } + connectionID := sqlConnectionIdentity(conn) + if schemaErr := s.setSchema(ctx, conn, connectionID, schema); schemaErr != nil { + _ = conn.Close() + if attempt == 0 && errors.Is(schemaErr, driver.ErrBadConn) { + s.resetSchemaCache() + continue + } + return nil, schemaErr + } + return conn, nil + } + return nil, driver.ErrBadConn +} + +func (s *server) setSchema(ctx context.Context, conn *sql.Conn, connectionID uintptr, schema string) error { + schema = strings.TrimSpace(schema) + if connectionID != 0 && s.schemaInitialized && s.schemaConnectionID == connectionID && s.currentSchema == schema { + return nil + } + if connectionID != 0 && !s.schemaInitialized && schema == "" { + s.currentSchema = "" + s.schemaInitialized = true + s.schemaConnectionID = connectionID + return nil + } + statement := "RESET search_path" + if schema != "" { + // Vastbase implicitly prioritizes its system catalog when it is not + // listed explicitly, matching the JDBC agent and DBeaver behavior. + statement = "SET search_path TO " + quoteIdentifier(schema) + } + if _, err := conn.ExecContext(ctx, statement); err != nil { + return err + } + s.currentSchema = schema + s.schemaInitialized = true + s.schemaConnectionID = connectionID + return nil +} + +func buildDSN(cp connectParams) string { + return buildDSNWithSSLMode(cp, agentInitialSSLMode(effectiveSSLMode(cp))) +} + +func buildDSNWithSSLMode(cp connectParams, sslMode string) string { + if value := strings.TrimSpace(cp.ConnectionString); value != "" && !isAgentJDBCURL(value) { + return rewriteNativeConnectionStringSSLMode(value, sslMode) + } + port := cp.Port + if port <= 0 { + port = agentDefaultPort + } + parts := []string{ + "host=" + quoteDSNValue(cp.Host), + fmt.Sprintf("port=%d", port), + "user=" + quoteDSNValue(cp.Username), + "password=" + quoteDSNValue(cp.Password), + "dbname=" + quoteDSNValue(cp.Database), + "sslmode=" + sslMode, + "connect_timeout=15", + } + if cp.CACertPath != "" { + parts = append(parts, "sslrootcert="+quoteDSNValue(cp.CACertPath)) + } + if cp.ClientCertPath != "" { + parts = append(parts, "sslcert="+quoteDSNValue(cp.ClientCertPath)) + } + if cp.ClientKeyPath != "" { + parts = append(parts, "sslkey="+quoteDSNValue(cp.ClientKeyPath)) + } + for _, parameter := range nativeURLParams(cp.URLParams) { + if !strings.EqualFold(parameter.Key, "sslmode") { + parts = append(parts, parameter.Key+"="+quoteDSNValue(parameter.Value)) + } + } + return strings.Join(parts, " ") +} + +func effectiveSSLMode(cp connectParams) string { + if value := strings.TrimSpace(cp.ConnectionString); value != "" && !isAgentJDBCURL(value) { + if sslMode, ok := nativeConnectionStringSSLMode(value); ok && sslMode != "" { + return sslMode + } + return "prefer" + } + sslMode := "" + for _, parameter := range nativeURLParams(cp.URLParams) { + if strings.EqualFold(parameter.Key, "sslmode") { + sslMode = strings.ToLower(strings.TrimSpace(parameter.Value)) + } + } + if sslMode != "" { + return sslMode + } + if cp.SSL { + return "verify-full" + } + return "prefer" +} + +func nativeConnectionStringSSLMode(value string) (string, bool) { + if isAgentNativeURL(value) { + query := value + if _, after, ok := strings.Cut(query, "?"); ok { + query = after + } else { + return "", false + } + query, _, _ = strings.Cut(query, "#") + sslMode := "" + found := false + for _, pair := range strings.Split(query, "&") { + key, rawValue, ok := strings.Cut(pair, "=") + if !ok { + continue + } + decodedKey, err := url.QueryUnescape(key) + if err != nil || !strings.EqualFold(decodedKey, "sslmode") { + continue + } + decodedValue, err := url.QueryUnescape(rawValue) + if err != nil { + decodedValue = rawValue + } + sslMode = strings.ToLower(strings.TrimSpace(decodedValue)) + found = true + } + return sslMode, found + } + + sslMode := "" + found := false + for _, field := range splitNativeDSNFields(value) { + key, rawValue, ok := strings.Cut(field, "=") + if !ok || !strings.EqualFold(strings.TrimSpace(key), "sslmode") { + continue + } + sslMode = strings.ToLower(unquoteNativeDSNValue(rawValue)) + found = true + } + return sslMode, found +} + +func rewriteNativeConnectionStringSSLMode(value, sslMode string) string { + if isAgentNativeURL(value) { + baseAndQuery, fragment, hasFragment := strings.Cut(value, "#") + base, query, hasQuery := strings.Cut(baseAndQuery, "?") + pairs := make([]string, 0) + if hasQuery { + for _, pair := range strings.Split(query, "&") { + key, _, _ := strings.Cut(pair, "=") + decodedKey, err := url.QueryUnescape(key) + if err == nil && strings.EqualFold(decodedKey, "sslmode") { + continue + } + if pair != "" { + pairs = append(pairs, pair) + } + } + } + pairs = append(pairs, "sslmode="+url.QueryEscape(sslMode)) + result := base + "?" + strings.Join(pairs, "&") + if hasFragment { + result += "#" + fragment + } + return result + } + + fields := splitNativeDSNFields(value) + result := make([]string, 0, len(fields)+1) + for _, field := range fields { + key, _, ok := strings.Cut(field, "=") + if ok && strings.EqualFold(strings.TrimSpace(key), "sslmode") { + continue + } + result = append(result, field) + } + result = append(result, "sslmode="+sslMode) + return strings.Join(result, " ") +} + +func splitNativeDSNFields(value string) []string { + fields := make([]string, 0) + for index := 0; index < len(value); { + for index < len(value) && isNativeDSNSpace(value[index]) { + index++ + } + if index >= len(value) { + break + } + start := index + for index < len(value) && value[index] != '=' { + index++ + } + if index >= len(value) { + fields = append(fields, strings.TrimSpace(value[start:])) + break + } + index++ + for index < len(value) && isNativeDSNSpace(value[index]) { + index++ + } + quoted := index < len(value) && value[index] == '\'' + if quoted { + index++ + } + for index < len(value) { + if value[index] == '\\' && index+1 < len(value) { + index += 2 + continue + } + if quoted { + if value[index] == '\'' { + index++ + break + } + } else if isNativeDSNSpace(value[index]) { + break + } + index++ + } + for index < len(value) && !isNativeDSNSpace(value[index]) { + index++ + } + fields = append(fields, strings.TrimSpace(value[start:index])) + } + return fields +} + +func unquoteNativeDSNValue(value string) string { + value = strings.TrimSpace(value) + if len(value) >= 2 && value[0] == '\'' && value[len(value)-1] == '\'' { + value = value[1 : len(value)-1] + } + return strings.TrimSpace(value) +} + +func isNativeDSNSpace(value byte) bool { + return value == ' ' || value == '\t' || value == '\n' || value == '\r' || value == '\f' +} + +func quoteDSNValue(value string) string { + return "'" + strings.ReplaceAll(strings.ReplaceAll(value, `\`, `\\`), "'", `\'`) + "'" +} + +func isSafeParamKey(value string) bool { + value = strings.TrimSpace(value) + if value == "" { + return false + } + for _, char := range value { + if !(char == '_' || char >= 'a' && char <= 'z' || char >= 'A' && char <= 'Z' || char >= '0' && char <= '9') { + return false + } + } + return true +} + +func normalizeValue(value any) any { + switch typed := value.(type) { + case nil: + return nil + case []byte: + if isTextBytes(typed) { + return string(typed) + } + return map[string]string{"$binary": base64.StdEncoding.EncodeToString(typed)} + case time.Time: + return typed.Format(time.RFC3339Nano) + case int8: + return int64(typed) + case int16: + return int64(typed) + case int32: + return int64(typed) + case float32: + return float64(typed) + default: + return typed + } +} + +func isTextBytes(value []byte) bool { + for _, char := range value { + if char == 0 || char < 0x09 || char > 0x0d && char < 0x20 { + return false + } + } + return true +} + +func decodeParams(params map[string]json.RawMessage, target any) error { + data, err := json.Marshal(params) + if err != nil { + return err + } + return json.Unmarshal(data, target) +} + +func stringParam(params map[string]json.RawMessage, key string) string { + var value string + _ = json.Unmarshal(params[key], &value) + return value +} + +func intParam(params map[string]json.RawMessage, key string) int { + var value int + _ = json.Unmarshal(params[key], &value) + return value +} + +func boolParam(params map[string]json.RawMessage, key string) bool { + var value bool + if raw, ok := params[key]; ok { + _ = json.Unmarshal(raw, &value) + } + return value +} + +func stringSliceParam(params map[string]json.RawMessage, key string) []string { + var values []string + if json.Unmarshal(params[key], &values) == nil { + return values + } + return nil +} + +func metadataListConstraintsFromParams(params map[string]json.RawMessage) metadataListConstraints { + return metadataListConstraints{ + Filter: stringParam(params, "filter"), + Limit: intParam(params, "limit"), + Offset: intParam(params, "offset"), + ObjectTypes: stringSliceParam(params, "object_types"), + } +} + +func queryOptionsFromParams(params map[string]json.RawMessage) queryOptions { + return queryOptions{ + SQL: stringParam(params, "sql"), + Database: stringParam(params, "database"), + Schema: stringParam(params, "schema"), + MaxRows: intParam(params, "maxRows"), + FetchSize: intParam(params, "fetchSize"), + TimeoutSecs: intParam(params, "timeoutSecs"), + } +} + +func errorResponse(id json.RawMessage, method, agentSessionID string, err error) response { + return response{JSONRPC: "2.0", ID: id, Error: classifyRPCError(method, agentSessionID, err)} +} + +func trimStatementSQL(sqlText string) string { + return strings.TrimRight(strings.TrimSpace(sqlText), "; \t\r\n") +} + +func isQuerySQL(sqlText string) bool { + lower := strings.ToLower(strings.TrimSpace(sqlText)) + return strings.HasPrefix(lower, "select") || strings.HasPrefix(lower, "with") || strings.HasPrefix(lower, "show") || strings.HasPrefix(lower, "explain") +} + +func quoteIdentifier(value string) string { + return `"` + strings.ReplaceAll(value, `"`, `""`) + `"` +} + +func quoteLiteral(value string) string { + return "'" + strings.ReplaceAll(value, "'", "''") + "'" +} + +func stringPtr(value string) *string { + if value == "" { + return nil + } + return &value +} diff --git a/agents/drivers/vastbase-go/main_test.go b/agents/drivers/vastbase-go/main_test.go new file mode 100644 index 000000000..4603be753 --- /dev/null +++ b/agents/drivers/vastbase-go/main_test.go @@ -0,0 +1,281 @@ +package main + +import ( + "context" + "database/sql" + "database/sql/driver" + "encoding/json" + "strings" + "sync" + "sync/atomic" + "testing" + + pq "gitcode.com/opengauss/openGauss-connector-go-pq" +) + +func TestVastbaseHandshakeAdvertisesMultiSessionSQLAgent(t *testing.T) { + runtime := &runtimeServer{sessions: map[string]*agentSession{}} + result, shutdown, err := runtime.dispatch("handshake", nil) + if err != nil { + t.Fatalf("handshake failed: %v", err) + } + if shutdown { + t.Fatal("handshake must not request shutdown") + } + payload, err := json.Marshal(result) + if err != nil { + t.Fatalf("marshal handshake: %v", err) + } + text := string(payload) + for _, expected := range []string{`"protocolVersion":2`, `"multi_session"`, `"metadata"`, `"paged_query"`, `"structured_error_v1"`} { + if !strings.Contains(text, expected) { + t.Fatalf("handshake missing %s: %s", expected, text) + } + } +} + +func TestQueryOptionsFromParams(t *testing.T) { + params := map[string]json.RawMessage{ + "sql": json.RawMessage(`"SELECT 1"`), + "database": json.RawMessage(`"dbx"`), + "schema": json.RawMessage(`"public"`), + "maxRows": json.RawMessage(`1000`), + "fetchSize": json.RawMessage(`250`), + "timeoutSecs": json.RawMessage(`15`), + } + expected := queryOptions{SQL: "SELECT 1", Database: "dbx", Schema: "public", MaxRows: 1000, FetchSize: 250, TimeoutSecs: 15} + if actual := queryOptionsFromParams(params); actual != expected { + t.Fatalf("queryOptionsFromParams() = %+v, want %+v", actual, expected) + } +} + +func TestVastbaseBuildDSNUsesNativeDefaultsForJDBCURL(t *testing.T) { + dsn := buildDSN(connectParams{ + Host: "vastbase.example.com", + Database: "postgres", + Username: "vbadmin", + Password: "secret", + ConnectionString: "jdbc:vastbase://vastbase.example.com:5432/postgres", + URLParams: "application_name=dbx", + }) + for _, expected := range []string{ + "host='vastbase.example.com'", + "port=5432", + "user='vbadmin'", + "password='secret'", + "dbname='postgres'", + "sslmode=prefer", + "application_name='dbx'", + } { + if !strings.Contains(dsn, expected) { + t.Fatalf("DSN missing %s: %s", expected, dsn) + } + } +} + +func TestVastbaseBuildDSNPreservesNativeConnectionString(t *testing.T) { + dsn := buildDSNWithSSLMode(connectParams{ + ConnectionString: "postgresql://vbadmin:secret@vastbase.example.com:5432/postgres?application_name=dbx&sslmode=disable", + }, "verify-full") + if !strings.Contains(dsn, "application_name=dbx") || !strings.Contains(dsn, "sslmode=verify-full") { + t.Fatalf("unexpected rewritten native DSN: %s", dsn) + } + if strings.Contains(dsn, "sslmode=disable") { + t.Fatalf("old sslmode was not replaced: %s", dsn) + } +} + +func TestVastbaseBuildDSNTranslatesJDBCParameters(t *testing.T) { + dsn := buildDSN(connectParams{ + Host: "vastbase.example.com", + Database: "postgres", + Username: "vbadmin", + Password: "secret", + URLParams: "targetServerType=master&connectTimeout=7¤tSchema=app&applicationName=dbx&sslmode=enable&autosave=always&enable_ce=1&db_compatibility=PG", + }) + for _, expected := range []string{ + "target_session_attrs='primary'", + "connect_timeout='7'", + "search_path='app'", + "application_name='dbx'", + "sslmode=require", + } { + if !strings.Contains(dsn, expected) { + t.Fatalf("translated DSN missing %s: %s", expected, dsn) + } + } + for _, rejected := range []string{"targetServerType", "currentSchema", "applicationName", "autosave", "enable_ce", "db_compatibility"} { + if strings.Contains(dsn, rejected) { + t.Fatalf("JDBC-only parameter leaked into DSN: %s", dsn) + } + } +} + +func TestVastbaseDriverRegistration(t *testing.T) { + if !containsString(sql.Drivers(), agentSQLDriverName) { + t.Fatalf("%s driver is not registered: %v", agentSQLDriverName, sql.Drivers()) + } +} + +func TestVastbaseObjectSourceNormalization(t *testing.T) { + tests := map[string]string{ + `(1,"CREATE FUNCTION f() RETURNS int AS ''SELECT 1'';")`: `CREATE FUNCTION f() RETURNS int AS ''SELECT 1'';`, + `("CREATE VIEW v AS SELECT 1")`: `CREATE VIEW v AS SELECT 1`, + `CREATE VIEW v AS SELECT 1`: `CREATE VIEW v AS SELECT 1`, + } + for input, expected := range tests { + if actual := normalizeAgentObjectSource(input); actual != expected { + t.Fatalf("normalizeAgentObjectSource(%q) = %q, want %q", input, actual, expected) + } + } +} + +func TestVastbaseDataTypesIncludeVectorFamilies(t *testing.T) { + types := agentDataTypes() + for _, expected := range []string{"floatvector", "halfvector", "int8vector", "sparsevector"} { + if !containsString(types, expected) { + t.Fatalf("missing Vastbase data type %s: %v", expected, types) + } + } +} + +func TestVastbaseMetadataErrorClassificationUsesOpenGaussCodes(t *testing.T) { + undefinedColumn := &pq.Error{Code: pq.ErrorCode("42703"), Message: "column a.attidentity does not exist"} + if !isUndefinedColumn(undefinedColumn, "attidentity") { + t.Fatal("undefined Vastbase column was not recognized") + } + undefinedFunction := &pq.Error{Code: pq.ErrorCode("42883"), Message: "function pg_get_expr does not exist"} + if !isUndefinedFunction(undefinedFunction, "pg_get_expr") { + t.Fatal("undefined Vastbase function was not recognized") + } +} + +func TestVastbaseModeUsesPostgresCatalog(t *testing.T) { + mode := detectAgentMode(nil, false) + if mode.compatibilityMode != "postgres" || !mode.postgresCatalog || mode.mysqlCompat { + t.Fatalf("unexpected default Vastbase mode: %+v", mode) + } + mysqlMode := detectAgentMode(nil, true) + if mysqlMode.compatibilityMode != "mysql" || !mysqlMode.postgresCatalog || !mysqlMode.mysqlCompat { + t.Fatalf("unexpected MySQL-compatible Vastbase mode: %+v", mysqlMode) + } +} + +func TestSchemaConnectionRecoversAfterCanceledDriverConnection(t *testing.T) { + registerVastbaseSchemaRetryDriver.Do(func() { + sql.Register("vastbase-schema-retry-test", &schemaRetryDriver{}) + }) + schemaRetryOpens.Store(0) + db, err := sql.Open("vastbase-schema-retry-test", "") + if err != nil { + t.Fatal(err) + } + defer db.Close() + db.SetMaxOpenConns(1) + db.SetMaxIdleConns(1) + + server := newServer() + server.db = db + conn, err := server.schemaConn(context.Background(), "public") + if err != nil { + t.Fatalf("schema connection did not recover: %v", err) + } + defer conn.Close() + if opens := schemaRetryOpens.Load(); opens != 2 { + t.Fatalf("expected one replacement connection, opened %d", opens) + } +} + +func TestValidateConnectionRecoversAfterCanceledDriverConnection(t *testing.T) { + registerVastbasePingRetryDriver.Do(func() { + sql.Register("vastbase-ping-retry-test", &pingRetryDriver{}) + }) + pingRetryOpens.Store(0) + db, err := sql.Open("vastbase-ping-retry-test", "") + if err != nil { + t.Fatal(err) + } + defer db.Close() + db.SetMaxOpenConns(1) + db.SetMaxIdleConns(1) + + server := newServer() + server.db = db + if err := server.validateConnection(); err != nil { + t.Fatalf("connection validation did not recover: %v", err) + } + if opens := pingRetryOpens.Load(); opens != 2 { + t.Fatalf("expected one replacement connection, opened %d", opens) + } +} + +func TestDisconnectResetsInformationSchemaCapabilityCache(t *testing.T) { + server := newServer() + server.infoColumnTypeUnsupported = true + server.infoUdtNameUnsupported = true + + if err := server.disconnect(); err != nil { + t.Fatal(err) + } + if server.infoColumnTypeUnsupported || server.infoUdtNameUnsupported { + t.Fatal("disconnect must reset cached information_schema capabilities") + } +} + +var ( + registerVastbaseSchemaRetryDriver sync.Once + registerVastbasePingRetryDriver sync.Once + schemaRetryOpens atomic.Int32 + pingRetryOpens atomic.Int32 +) + +type schemaRetryDriver struct{} + +func (*schemaRetryDriver) Open(string) (driver.Conn, error) { + return &schemaRetryConn{bad: schemaRetryOpens.Add(1) == 1}, nil +} + +type schemaRetryConn struct { + bad bool +} + +func (*schemaRetryConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip } +func (*schemaRetryConn) Close() error { return nil } +func (*schemaRetryConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip } + +func (conn *schemaRetryConn) ExecContext(context.Context, string, []driver.NamedValue) (driver.Result, error) { + if conn.bad { + return nil, driver.ErrBadConn + } + return driver.RowsAffected(0), nil +} + +type pingRetryDriver struct{} + +func (*pingRetryDriver) Open(string) (driver.Conn, error) { + return &pingRetryConn{bad: pingRetryOpens.Add(1) == 1}, nil +} + +type pingRetryConn struct { + bad bool +} + +func (*pingRetryConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip } +func (*pingRetryConn) Close() error { return nil } +func (*pingRetryConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip } + +func (conn *pingRetryConn) Ping(context.Context) error { + if conn.bad { + return driver.ErrBadConn + } + return nil +} + +func containsString(values []string, expected string) bool { + for _, value := range values { + if value == expected { + return true + } + } + return false +} diff --git a/agents/drivers/vastbase-go/protocol_error.go b/agents/drivers/vastbase-go/protocol_error.go new file mode 100644 index 000000000..53c527637 --- /dev/null +++ b/agents/drivers/vastbase-go/protocol_error.go @@ -0,0 +1,154 @@ +package main + +import ( + "context" + "database/sql/driver" + "errors" + "fmt" + "io" + "net" + "strings" + + pq "gitcode.com/opengauss/openGauss-connector-go-pq" +) + +type rpcError struct { + Code int `json:"code"` + Message string `json:"message"` + Data *rpcErrorData `json:"data,omitempty"` +} + +type rpcErrorData struct { + Category string `json:"category"` + Retryable bool `json:"retryable"` + SessionDisposition string `json:"sessionDisposition"` + Stage string `json:"stage"` + ContractVersion int `json:"contractVersion"` + OperationOutcome string `json:"operationOutcome"` + SQLState string `json:"sqlState,omitempty"` + ExceptionClass string `json:"exceptionClass,omitempty"` + AgentSessionID string `json:"agentSessionId,omitempty"` +} + +func classifyRPCError(method, agentSessionID string, err error) *rpcError { + stage := rpcErrorStage(method) + data := &rpcErrorData{ + Category: "protocol", + Retryable: false, + SessionDisposition: "keep", + Stage: stage, + ContractVersion: 1, + OperationOutcome: rpcOperationOutcome(stage), + ExceptionClass: safeRPCDiagnostic(fmt.Sprintf("%T", err), 160), + AgentSessionID: strings.TrimSpace(agentSessionID), + } + if errors.Is(err, errOperationCapacity) { + data.Category = "resource" + data.Retryable = true + return &rpcError{Code: -1, Message: err.Error(), Data: data} + } + + var databaseError *pq.Error + if errors.As(err, &databaseError) { + data.SQLState = safeRPCDiagnostic(databaseError.SQLState(), 16) + switch { + case databaseError.Code == pq.ErrorCode("57014"): + data.Category = "canceled" + data.SessionDisposition = "quarantine" + case stage == "connect" || stage == "validate" || strings.HasPrefix(databaseError.SQLState(), "08"): + data.Category = "connection" + data.Retryable = stage == "connect" || stage == "validate" + if stage != "connect" { + data.SessionDisposition = "quarantine" + } + default: + data.Category = "sql" + } + } else if errors.Is(err, context.Canceled) { + data.Category = "canceled" + data.SessionDisposition = "quarantine" + } else if errors.Is(err, context.DeadlineExceeded) || isTimeoutError(err) { + data.Category = "timeout" + data.SessionDisposition = "quarantine" + } else if isConnectionError(err) { + data.Category = "connection" + data.Retryable = stage == "connect" || stage == "validate" + if stage != "connect" { + data.SessionDisposition = "quarantine" + } + } + + return &rpcError{Code: -1, Message: err.Error(), Data: data} +} + +func rpcErrorStage(method string) string { + switch method { + case "connect", "open_session", "test_connection": + return "connect" + case "validate_connection", "validate_session": + return "validate" + case "cancel_session": + return "cancel" + case "close_session", "disconnect", "close_query_session", "close_table_read_session", "shutdown": + return "close" + case "fetch_query_page", "fetch_table_read_page": + return "fetch" + case "handshake", "": + return "request" + default: + return "execute" + } +} + +func rpcOperationOutcome(stage string) string { + switch stage { + case "request", "connect", "validate": + return "not_started" + default: + return "unknown" + } +} + +func isTimeoutError(err error) bool { + var timeout interface{ Timeout() bool } + return errors.As(err, &timeout) && timeout.Timeout() +} + +func isConnectionError(err error) bool { + if errors.Is(err, driver.ErrBadConn) || errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) { + return true + } + var networkError *net.OpError + if errors.As(err, &networkError) { + return true + } + lower := strings.ToLower(err.Error()) + for _, marker := range []string{ + "connection refused", + "connection reset", + "broken pipe", + "connection closed", + "connection lost", + "driver: bad connection", + "unexpected eof", + "no route to host", + } { + if strings.Contains(lower, marker) { + return true + } + } + return false +} + +func safeRPCDiagnostic(value string, maxLength int) string { + var result strings.Builder + for _, char := range value { + if result.Len() >= maxLength { + break + } + if char >= 0x21 && char <= 0x7e { + result.WriteRune(char) + } + } + return result.String() +} diff --git a/agents/drivers/vastbase-go/protocol_error_test.go b/agents/drivers/vastbase-go/protocol_error_test.go new file mode 100644 index 000000000..2f071ee32 --- /dev/null +++ b/agents/drivers/vastbase-go/protocol_error_test.go @@ -0,0 +1,62 @@ +package main + +import ( + "context" + "encoding/json" + "testing" + + pq "gitcode.com/opengauss/openGauss-connector-go-pq" +) + +func TestStructuredRPCErrorClassification(t *testing.T) { + tests := []struct { + name string + method string + err error + category string + retryable bool + disposition string + sqlState string + }{ + {name: "sql", method: "execute_query", err: &pq.Error{Code: pq.ErrorCode("42P01"), Message: "relation missing"}, category: "sql", disposition: "keep", sqlState: "42P01"}, + {name: "connection", method: "execute_query", err: &pq.Error{Code: pq.ErrorCode("08006"), Message: "connection failure"}, category: "connection", disposition: "quarantine", sqlState: "08006"}, + {name: "connect", method: "open_session", err: &pq.Error{Code: pq.ErrorCode("28P01"), Message: "bad password"}, category: "connection", retryable: true, disposition: "keep", sqlState: "28P01"}, + {name: "query canceled", method: "execute_query", err: &pq.Error{Code: pq.ErrorCode("57014"), Message: "canceling statement"}, category: "canceled", disposition: "quarantine", sqlState: "57014"}, + {name: "context canceled", method: "execute_query", err: context.Canceled, category: "canceled", disposition: "quarantine"}, + {name: "timeout", method: "execute_query", err: context.DeadlineExceeded, category: "timeout", disposition: "quarantine"}, + {name: "capacity", method: "execute_query", err: errOperationCapacity, category: "resource", retryable: true, disposition: "keep"}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + rpcErr := classifyRPCError(test.method, "session-1", test.err) + if rpcErr.Data.Category != test.category || rpcErr.Data.Retryable != test.retryable || rpcErr.Data.SessionDisposition != test.disposition || rpcErr.Data.SQLState != test.sqlState { + t.Fatalf("unexpected classification: %+v", rpcErr.Data) + } + if rpcErr.Data.AgentSessionID != "session-1" || rpcErr.Data.ContractVersion != 1 { + t.Fatalf("missing structured error identity: %+v", rpcErr.Data) + } + }) + } +} + +func TestStructuredRPCErrorContainsRequiredContractFields(t *testing.T) { + rpcErr := classifyRPCError("fetch_query_page", "session-2", context.DeadlineExceeded) + payload, err := json.Marshal(rpcErr) + if err != nil { + t.Fatal(err) + } + var decoded struct { + Data map[string]any `json:"data"` + } + if err := json.Unmarshal(payload, &decoded); err != nil { + t.Fatal(err) + } + for _, key := range []string{"category", "retryable", "sessionDisposition", "stage", "contractVersion", "operationOutcome", "agentSessionId"} { + if _, ok := decoded.Data[key]; !ok { + t.Fatalf("structured error missing %s: %s", key, payload) + } + } + if decoded.Data["stage"] != "fetch" || decoded.Data["operationOutcome"] != "unknown" { + t.Fatalf("unexpected fetch error contract: %s", payload) + } +} diff --git a/agents/drivers/vastbase-go/runtime_pool.go b/agents/drivers/vastbase-go/runtime_pool.go new file mode 100644 index 000000000..70b97aa74 --- /dev/null +++ b/agents/drivers/vastbase-go/runtime_pool.go @@ -0,0 +1,247 @@ +package main + +import ( + "context" + "crypto/sha256" + "database/sql" + "errors" + "fmt" + "os" + "strconv" + "strings" + "sync" + "time" +) + +const ( + defaultRuntimePoolSize = 32 + defaultRuntimeMetadataLimit = 8 + defaultValidatorPoolSize = 8 + connectionRuntimeGracePeriod = 30 * time.Second + operationPermitTimeout = 30 * time.Second +) + +var errOperationCapacity = errors.New("agent operation capacity is temporarily exhausted") + +type connectionRuntime struct { + mu sync.Mutex + validator *sql.DB + listTablesStatement *sql.Stmt + permits chan struct{} + metadataPermits chan struct{} + references int + lastReleased time.Time +} + +func newConnectionRuntime() *connectionRuntime { + poolSize := runtimePoolSize() + return &connectionRuntime{ + permits: make(chan struct{}, poolSize), + metadataPermits: make(chan struct{}, runtimeMetadataLimit(poolSize)), + } +} + +func runtimeMetadataLimit(poolSize int) int { + value := min(defaultRuntimeMetadataLimit, poolSize) + if raw := os.Getenv("DBX_AGENT_VASTBASE_MAX_CONCURRENT_METADATA"); raw != "" { + if parsed, err := strconv.Atoi(raw); err == nil && parsed >= 1 && parsed <= poolSize { + value = parsed + } + } + return value +} + +func runtimePoolSize() int { + value := defaultRuntimePoolSize + if raw := os.Getenv("DBX_AGENT_VASTBASE_MAX_CONCURRENT_OPERATIONS"); raw != "" { + if parsed, err := strconv.Atoi(raw); err == nil && parsed >= 1 && parsed <= 32 { + value = parsed + } + } + return value +} + +func (connectionRuntime *connectionRuntime) validate(cp connectParams, opener agentDBOpener) error { + connectionRuntime.mu.Lock() + validator := connectionRuntime.validator + if validator == nil { + db, err := openAndPingDB(cp, defaultConnectTimeout, opener) + if err != nil { + connectionRuntime.mu.Unlock() + return err + } + poolSize := cap(connectionRuntime.metadataPermits) + db.SetMaxOpenConns(poolSize) + db.SetMaxIdleConns(poolSize) + db.SetConnMaxLifetime(5 * time.Minute) + connectionRuntime.validator = db + connectionRuntime.mu.Unlock() + return nil + } + connectionRuntime.mu.Unlock() + + ctx, cancel := context.WithTimeout(context.Background(), defaultConnectTimeout) + defer cancel() + return validator.PingContext(ctx) +} + +func (connectionRuntime *connectionRuntime) acquire(metadata bool) (func(), error) { + ctx, cancel := context.WithTimeout(context.Background(), operationPermitTimeout) + defer cancel() + metadataAcquired := false + if metadata { + select { + case connectionRuntime.metadataPermits <- struct{}{}: + metadataAcquired = true + case <-ctx.Done(): + return nil, errOperationCapacity + } + } + select { + case connectionRuntime.permits <- struct{}{}: + return func() { + <-connectionRuntime.permits + if metadataAcquired { + <-connectionRuntime.metadataPermits + } + }, nil + case <-ctx.Done(): + if metadataAcquired { + <-connectionRuntime.metadataPermits + } + return nil, errOperationCapacity + } +} + +func (connectionRuntime *connectionRuntime) close() error { + connectionRuntime.mu.Lock() + validator := connectionRuntime.validator + listTablesStatement := connectionRuntime.listTablesStatement + connectionRuntime.validator = nil + connectionRuntime.listTablesStatement = nil + connectionRuntime.mu.Unlock() + if listTablesStatement != nil { + _ = listTablesStatement.Close() + } + if validator == nil { + return nil + } + return validator.Close() +} + +func (connectionRuntime *connectionRuntime) database() *sql.DB { + connectionRuntime.mu.Lock() + defer connectionRuntime.mu.Unlock() + return connectionRuntime.validator +} + +func (connectionRuntime *connectionRuntime) queryListTables(query, schema string) (*sql.Rows, error) { + connectionRuntime.mu.Lock() + statement := connectionRuntime.listTablesStatement + if statement == nil { + if connectionRuntime.validator == nil { + connectionRuntime.mu.Unlock() + return nil, errors.New("connection runtime is not initialized") + } + prepared, err := connectionRuntime.validator.Prepare(query) + if err != nil { + connectionRuntime.mu.Unlock() + return nil, err + } + connectionRuntime.listTablesStatement = prepared + statement = prepared + } + connectionRuntime.mu.Unlock() + return statement.Query(schema) +} + +func (s *server) acquireOperationPermit(method string) (func(), error) { + if s.connectionRuntime == nil { + return func() {}, nil + } + metadata := isMetadataOperation(method) || strings.EqualFold(strings.TrimSpace(s.params.SessionRole), "metadata") + return s.connectionRuntime.acquire(metadata) +} + +func (s *server) metadataDatabase() (*sql.DB, error) { + if s.connectionRuntime != nil && !s.sessionAffinity { + if db := s.connectionRuntime.database(); db != nil { + return db, nil + } + } + return s.requireDB() +} + +func isMetadataOperation(method string) bool { + switch method { + case "connection_info", "list_databases", "list_schemas", "list_tables", "get_table_comment", "list_objects", + "list_data_types", "completion_assistant_search_v1", "get_columns", "list_indexes", "list_foreign_keys", + "list_triggers", "get_object_source", "get_table_ddl", "get_explain_info": + return true + default: + return false + } +} + +func (r *runtimeServer) acquireConnectionRuntime(cp connectParams) (*connectionRuntime, string) { + key := connectionRuntimeKey(cp) + r.connectionRuntimeMu.Lock() + if r.connectionRuntimes == nil { + r.connectionRuntimes = map[string]*connectionRuntime{} + } + r.closeExpiredConnectionRuntimesLocked(time.Now()) + connectionRuntime := r.connectionRuntimes[key] + if connectionRuntime == nil { + connectionRuntime = newConnectionRuntime() + r.connectionRuntimes[key] = connectionRuntime + } + connectionRuntime.references++ + r.connectionRuntimeMu.Unlock() + + return connectionRuntime, key +} + +func (r *runtimeServer) releaseConnectionRuntime(key string) { + if key == "" { + return + } + r.connectionRuntimeMu.Lock() + if connectionRuntime := r.connectionRuntimes[key]; connectionRuntime != nil { + if connectionRuntime.references > 0 { + connectionRuntime.references-- + } + if connectionRuntime.references == 0 { + connectionRuntime.lastReleased = time.Now() + } + } + r.connectionRuntimeMu.Unlock() +} + +func (r *runtimeServer) closeExpiredConnectionRuntimesLocked(now time.Time) { + for key, connectionRuntime := range r.connectionRuntimes { + if connectionRuntime.references == 0 && !connectionRuntime.lastReleased.IsZero() && now.Sub(connectionRuntime.lastReleased) >= connectionRuntimeGracePeriod { + _ = connectionRuntime.close() + delete(r.connectionRuntimes, key) + } + } +} + +func (r *runtimeServer) closeConnectionRuntimes() error { + r.connectionRuntimeMu.Lock() + runtimes := r.connectionRuntimes + r.connectionRuntimes = map[string]*connectionRuntime{} + r.connectionRuntimeMu.Unlock() + var firstErr error + for _, connectionRuntime := range runtimes { + if err := connectionRuntime.close(); err != nil && firstErr == nil { + firstErr = err + } + } + return firstErr +} + +func connectionRuntimeKey(cp connectParams) string { + identity := fmt.Sprintf("%s\x00mysql=%t", buildDSNWithSSLMode(cp, agentInitialSSLMode(effectiveSSLMode(cp))), cp.MySQLCompatMode) + digest := sha256.Sum256([]byte(identity)) + return fmt.Sprintf("%x", digest[:]) +} diff --git a/agents/drivers/vastbase-go/runtime_pool_test.go b/agents/drivers/vastbase-go/runtime_pool_test.go new file mode 100644 index 000000000..215b8c636 --- /dev/null +++ b/agents/drivers/vastbase-go/runtime_pool_test.go @@ -0,0 +1,258 @@ +package main + +import ( + "context" + "database/sql" + "database/sql/driver" + "fmt" + "io" + "sync" + "sync/atomic" + "testing" + "time" +) + +func TestConnectionRuntimeReusesAuthenticatedValidator(t *testing.T) { + state := &runtimePoolTestState{} + opener := runtimePoolTestOpener(t, state) + connectionRuntime := newConnectionRuntime() + t.Cleanup(func() { _ = connectionRuntime.close() }) + + if err := connectionRuntime.validate(connectParams{}, opener); err != nil { + t.Fatal(err) + } + if err := connectionRuntime.validate(connectParams{}, opener); err != nil { + t.Fatal(err) + } + if opens := state.openCount(); opens != 1 { + t.Fatalf("validator did not reuse its authenticated connection: opened %d", opens) + } + if pings := state.pingCount(); pings != 2 { + t.Fatalf("unexpected validator ping count: %d", pings) + } +} + +func TestConnectionRuntimeRetainsMetadataPoolConnections(t *testing.T) { + state := &runtimePoolTestState{} + connectionRuntime := newConnectionRuntime() + t.Cleanup(func() { _ = connectionRuntime.close() }) + if err := connectionRuntime.validate(connectParams{}, runtimePoolTestOpener(t, state)); err != nil { + t.Fatal(err) + } + + acquirePool := func() { + connections := make([]*sql.Conn, 0, defaultValidatorPoolSize) + for range defaultValidatorPoolSize { + conn, err := connectionRuntime.database().Conn(context.Background()) + if err != nil { + t.Fatal(err) + } + connections = append(connections, conn) + } + for _, conn := range connections { + if err := conn.Close(); err != nil { + t.Fatal(err) + } + } + } + + acquirePool() + if opens := state.openCount(); opens != defaultValidatorPoolSize { + t.Fatalf("metadata pool opened %d connections, want %d", opens, defaultValidatorPoolSize) + } + acquirePool() + if opens := state.openCount(); opens != defaultValidatorPoolSize { + t.Fatalf("metadata pool discarded idle connections and reopened %d total", opens) + } +} + +func TestConnectWithRuntimeUsesSharedMetadataUntilSessionAffinity(t *testing.T) { + state := &runtimePoolTestState{} + opener := runtimePoolTestOpener(t, state) + connectionRuntime := newConnectionRuntime() + t.Cleanup(func() { _ = connectionRuntime.close() }) + server := newServer() + server.openDatabase = opener + + if err := server.connectWithRuntime(connectParams{}, connectionRuntime); err != nil { + t.Fatal(err) + } + if opens := state.openCount(); opens != 1 { + t.Fatalf("logical connect opened a private physical connection: %d", opens) + } + if err := server.validateConnection(); err != nil { + t.Fatal(err) + } + if opens := state.openCount(); opens != 1 { + t.Fatalf("stateless metadata opened a private physical connection: %d", opens) + } + server.noteSQLSessionState("SET ROLE analyst") + if err := server.validateConnection(); err != nil { + t.Fatal(err) + } + if opens := state.openCount(); opens != 2 { + t.Fatalf("session-affine metadata did not open its private physical connection: %d", opens) + } + if err := server.disconnect(); err != nil { + t.Fatal(err) + } +} + +func TestConnectionRuntimeSharesListTablesStatementAcrossSessions(t *testing.T) { + state := &runtimePoolTestState{} + opener := runtimePoolTestOpener(t, state) + connectionRuntime := newConnectionRuntime() + first := newServer() + first.openDatabase = opener + second := newServer() + second.openDatabase = opener + + if err := first.connectWithRuntime(connectParams{}, connectionRuntime); err != nil { + t.Fatal(err) + } + if err := second.connectWithRuntime(connectParams{}, connectionRuntime); err != nil { + t.Fatal(err) + } + for _, server := range []*server{first, second} { + rows, err := server.cachedListTablesQuery("SELECT value FROM tables WHERE schema = $1", "public") + if err != nil { + t.Fatal(err) + } + if err := rows.Close(); err != nil { + t.Fatal(err) + } + } + if prepares := state.prepareCount(); prepares != 1 { + t.Fatalf("shared list-tables statement prepared %d times, want 1", prepares) + } + if err := first.disconnect(); err != nil { + t.Fatal(err) + } + if err := second.disconnect(); err != nil { + t.Fatal(err) + } + if err := connectionRuntime.close(); err != nil { + t.Fatal(err) + } + if closes := state.statementCloseCount(); closes != 1 { + t.Fatalf("shared list-tables statement closed %d times, want 1", closes) + } +} + +func TestConnectionRuntimeLimitsConcurrentOperations(t *testing.T) { + t.Setenv("DBX_AGENT_VASTBASE_MAX_CONCURRENT_OPERATIONS", "2") + connectionRuntime := newConnectionRuntime() + var active atomic.Int32 + var peak atomic.Int32 + var waitGroup sync.WaitGroup + for range 8 { + waitGroup.Add(1) + go func() { + defer waitGroup.Done() + release, err := connectionRuntime.acquire(false) + if err != nil { + t.Errorf("acquire permit: %v", err) + return + } + current := active.Add(1) + for current > peak.Load() && !peak.CompareAndSwap(peak.Load(), current) { + } + time.Sleep(10 * time.Millisecond) + active.Add(-1) + release() + }() + } + waitGroup.Wait() + if value := peak.Load(); value != 2 { + t.Fatalf("operation concurrency peak = %d, want 2", value) + } +} + +func TestConnectionRuntimeKeySeparatesCredentialsWithoutExposingThem(t *testing.T) { + first := connectionRuntimeKey(connectParams{Host: "db", Database: "app", Username: "user", Password: "secret-a"}) + second := connectionRuntimeKey(connectParams{Host: "db", Database: "app", Username: "user", Password: "secret-b"}) + if first == second { + t.Fatal("different credentials shared one runtime key") + } + if len(first) != 64 || first == "secret-a" { + t.Fatalf("runtime key is not a SHA-256 digest: %q", first) + } +} + +var runtimePoolDriverSequence atomic.Uint64 + +type runtimePoolTestState struct { + opens atomic.Int32 + pings atomic.Int32 + prepares atomic.Int32 + statementCloses atomic.Int32 +} + +func (state *runtimePoolTestState) openCount() int32 { return state.opens.Load() } +func (state *runtimePoolTestState) pingCount() int32 { return state.pings.Load() } +func (state *runtimePoolTestState) prepareCount() int32 { return state.prepares.Load() } +func (state *runtimePoolTestState) statementCloseCount() int32 { return state.statementCloses.Load() } + +type runtimePoolTestDriver struct { + state *runtimePoolTestState +} + +func (testDriver *runtimePoolTestDriver) Open(string) (driver.Conn, error) { + testDriver.state.opens.Add(1) + return &runtimePoolTestConn{state: testDriver.state}, nil +} + +type runtimePoolTestConn struct { + state *runtimePoolTestState +} + +func (conn *runtimePoolTestConn) Prepare(string) (driver.Stmt, error) { + conn.state.prepares.Add(1) + return &runtimePoolTestStmt{state: conn.state}, nil +} + +func (*runtimePoolTestConn) Close() error { return nil } +func (*runtimePoolTestConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip } + +func (conn *runtimePoolTestConn) Ping(context.Context) error { + conn.state.pings.Add(1) + return nil +} + +type runtimePoolTestStmt struct { + state *runtimePoolTestState +} + +func (stmt *runtimePoolTestStmt) Close() error { + stmt.state.statementCloses.Add(1) + return nil +} + +func (*runtimePoolTestStmt) NumInput() int { return -1 } +func (*runtimePoolTestStmt) Exec([]driver.Value) (driver.Result, error) { + return driver.RowsAffected(0), nil +} +func (*runtimePoolTestStmt) Query([]driver.Value) (driver.Rows, error) { + return &runtimePoolTestRows{}, nil +} + +type runtimePoolTestRows struct{} + +func (*runtimePoolTestRows) Columns() []string { return []string{"value"} } +func (*runtimePoolTestRows) Close() error { return nil } +func (*runtimePoolTestRows) Next([]driver.Value) error { return io.EOF } + +func runtimePoolTestOpener(t *testing.T, state *runtimePoolTestState) agentDBOpener { + t.Helper() + driverName := fmt.Sprintf("vastbase-runtime-pool-%d", runtimePoolDriverSequence.Add(1)) + sql.Register(driverName, &runtimePoolTestDriver{state: state}) + return func(connectParams, string) (*sql.DB, error) { + db, err := sql.Open(driverName, "") + if err != nil { + return nil, err + } + db.SetMaxOpenConns(1) + db.SetMaxIdleConns(1) + return db, nil + } +} diff --git a/agents/drivers/vastbase-go/spatial.go b/agents/drivers/vastbase-go/spatial.go new file mode 100644 index 000000000..3ba897e88 --- /dev/null +++ b/agents/drivers/vastbase-go/spatial.go @@ -0,0 +1,536 @@ +package main + +import ( + "encoding/binary" + "encoding/hex" + "fmt" + "math" + "strconv" + "strings" + "unicode/utf8" +) + +type spatialColumn struct { + ColumnIndex int `json:"column_index"` + SRID *uint32 `json:"srid"` +} + +type spatialDecoder struct { + indices []int + spatial []bool + observed bool + sridByColumn map[int]uint32 +} + +func newSpatialDecoder(columnTypes []string) *spatialDecoder { + indices := make([]int, 0, len(columnTypes)) + for index, columnType := range columnTypes { + if isSpatialColumnType(columnType) { + indices = append(indices, index) + } + } + if len(indices) == 0 { + return nil + } + spatial := make([]bool, len(columnTypes)) + for _, index := range indices { + spatial[index] = true + } + return &spatialDecoder{indices: indices, spatial: spatial, sridByColumn: map[int]uint32{}} +} + +func isSpatialColumnType(columnType string) bool { + normalized := strings.ToLower(strings.TrimSpace(columnType)) + if index := strings.LastIndexByte(normalized, '.'); index >= 0 { + normalized = normalized[index+1:] + } + if index := strings.IndexByte(normalized, '('); index >= 0 { + normalized = normalized[:index] + } + return normalized == "geometry" || normalized == "geography" +} + +func (decoder *spatialDecoder) normalizeRow(values []any) ([]any, []*uint32, error) { + rowSRIDs := make([]*uint32, len(values)) + decoder.observed = true + for _, index := range decoder.indices { + if index >= len(values) { + continue + } + value, srid := decodeSpatialValue(values[index]) + values[index] = value + rowSRIDs[index] = srid + if srid != nil { + if _, exists := decoder.sridByColumn[index]; !exists { + decoder.sridByColumn[index] = *srid + } + } + } + for index, value := range values { + if index >= len(decoder.spatial) || !decoder.spatial[index] { + values[index] = normalizeValue(value) + } + } + return values, rowSRIDs, nil +} + +func (decoder *spatialDecoder) columns() []spatialColumn { + if decoder == nil || !decoder.observed { + return nil + } + columns := make([]spatialColumn, 0, len(decoder.indices)) + for _, index := range decoder.indices { + var srid *uint32 + if value, ok := decoder.sridByColumn[index]; ok { + srid = uint32Pointer(value) + } + columns = append(columns, spatialColumn{ColumnIndex: index, SRID: srid}) + } + return columns +} + +func spatialResultMetadata(decoder *spatialDecoder, values [][]*uint32) ([]spatialColumn, [][]*uint32) { + columns := decoder.columns() + if len(columns) == 0 { + return nil, nil + } + return columns, values +} + +func finishSpatialPage(result queryPageResult, decoder *spatialDecoder, values [][]*uint32) queryPageResult { + result.SpatialColumns, result.SpatialValues = spatialResultMetadata(decoder, values) + return result +} + +func decodeSpatialValue(value any) (any, *uint32) { + if value == nil { + return nil, nil + } + switch typed := value.(type) { + case []byte: + if decoded, ok := decodeWKBGeometry(typed); ok { + return decoded.WKT, decoded.SRID + } + if utf8.Valid(typed) { + return decodeSpatialText(string(typed)) + } + return "0x" + hex.EncodeToString(typed), nil + case string: + return decodeSpatialText(typed) + default: + return decodeSpatialText(fmt.Sprint(value)) + } +} + +func decodeSpatialText(value string) (any, *uint32) { + if len(value) >= 7 && strings.EqualFold(value[:5], "SRID=") { + if separator := strings.IndexByte(value[5:], ';'); separator >= 0 { + separator += 5 + if parsed, err := strconv.ParseInt(value[5:separator], 10, 32); err == nil { + var srid *uint32 + if parsed > 0 { + srid = uint32Pointer(uint32(parsed)) + } + return value[separator+1:], srid + } + } + } + raw, ok := parseSpatialHex(value) + if !ok { + return value, nil + } + if decoded, decodedOK := decodeWKBGeometry(raw); decodedOK { + return decoded.WKT, decoded.SRID + } + if strings.HasPrefix(value, "0x") || strings.HasPrefix(value, "0X") { + return value, nil + } + return "0x" + value, nil +} + +func parseSpatialHex(value string) ([]byte, bool) { + normalized := value + if strings.HasPrefix(normalized, "0x") || strings.HasPrefix(normalized, "0X") || strings.HasPrefix(normalized, `\x`) || strings.HasPrefix(normalized, `\X`) { + normalized = normalized[2:] + } + if len(normalized) < 10 || len(normalized)%2 != 0 || (normalized[:2] != "00" && normalized[:2] != "01") { + return nil, false + } + raw, err := hex.DecodeString(normalized) + return raw, err == nil +} + +func uint32Pointer(value uint32) *uint32 { + result := value + return &result +} + +type decodedWKBGeometry struct { + WKT string + SRID *uint32 +} + +type wkbDimensions struct { + hasZ bool + hasM bool +} + +func (dimensions wkbDimensions) suffix() string { + switch { + case dimensions.hasZ && dimensions.hasM: + return " ZM" + case dimensions.hasZ: + return " Z" + case dimensions.hasM: + return " M" + default: + return "" + } +} + +func (dimensions wkbDimensions) coordinateLength() int { + length := 2 + if dimensions.hasZ { + length++ + } + if dimensions.hasM { + length++ + } + return length +} + +type wkbGeometry struct { + kind uint32 + dimensions wkbDimensions + coords []float64 + points [][]float64 + rings [][][]float64 + multiPoint []wkbPoint + polygons [][][][]float64 + children []wkbGeometry +} + +type wkbPoint struct { + coords []float64 + empty bool +} + +func (geometry wkbGeometry) wkt() string { + suffix := geometry.dimensions.suffix() + switch geometry.kind { + case 1: + if geometry.coords == nil { + return "POINT" + suffix + " EMPTY" + } + return "POINT" + suffix + "(" + formatWKBCoordinate(geometry.coords) + ")" + case 2: + if len(geometry.points) == 0 { + return "LINESTRING" + suffix + " EMPTY" + } + return "LINESTRING" + suffix + "(" + formatWKBCoordinateSequence(geometry.points) + ")" + case 3: + if len(geometry.rings) == 0 { + return "POLYGON" + suffix + " EMPTY" + } + return "POLYGON" + suffix + "(" + formatWKBRings(geometry.rings) + ")" + case 4: + if len(geometry.multiPoint) == 0 { + return "MULTIPOINT" + suffix + " EMPTY" + } + parts := make([]string, len(geometry.multiPoint)) + for index, point := range geometry.multiPoint { + if point.empty { + parts[index] = "EMPTY" + } else { + parts[index] = "(" + formatWKBCoordinate(point.coords) + ")" + } + } + return "MULTIPOINT" + suffix + "(" + strings.Join(parts, ",") + ")" + case 5: + if len(geometry.rings) == 0 { + return "MULTILINESTRING" + suffix + " EMPTY" + } + return "MULTILINESTRING" + suffix + "(" + formatWKBRings(geometry.rings) + ")" + case 6: + if len(geometry.polygons) == 0 { + return "MULTIPOLYGON" + suffix + " EMPTY" + } + parts := make([]string, len(geometry.polygons)) + for index, polygon := range geometry.polygons { + parts[index] = "(" + formatWKBRings(polygon) + ")" + } + return "MULTIPOLYGON" + suffix + "(" + strings.Join(parts, ",") + ")" + case 7: + if len(geometry.children) == 0 { + return "GEOMETRYCOLLECTION" + suffix + " EMPTY" + } + parts := make([]string, len(geometry.children)) + for index, child := range geometry.children { + parts[index] = child.wkt() + } + return "GEOMETRYCOLLECTION" + suffix + "(" + strings.Join(parts, ",") + ")" + default: + return "" + } +} + +func formatWKBCoordinate(coordinates []float64) string { + values := make([]string, len(coordinates)) + for index, value := range coordinates { + switch { + case math.IsInf(value, 1): + values[index] = "inf" + case math.IsInf(value, -1): + values[index] = "-inf" + default: + values[index] = strconv.FormatFloat(value, 'g', -1, 64) + } + } + return strings.Join(values, " ") +} + +func formatWKBCoordinateSequence(points [][]float64) string { + values := make([]string, len(points)) + for index, point := range points { + values[index] = formatWKBCoordinate(point) + } + return strings.Join(values, ",") +} + +func formatWKBRings(rings [][][]float64) string { + values := make([]string, len(rings)) + for index, ring := range rings { + values[index] = "(" + formatWKBCoordinateSequence(ring) + ")" + } + return strings.Join(values, ",") +} + +type wkbReader struct { + raw []byte + position int +} + +func (reader *wkbReader) readByte() (byte, bool) { + if reader.position >= len(reader.raw) { + return 0, false + } + value := reader.raw[reader.position] + reader.position++ + return value, true +} + +func (reader *wkbReader) readUint32(order binary.ByteOrder) (uint32, bool) { + end := reader.position + 4 + if end > len(reader.raw) { + return 0, false + } + value := order.Uint32(reader.raw[reader.position:end]) + reader.position = end + return value, true +} + +func (reader *wkbReader) readFloat64(order binary.ByteOrder) (float64, bool) { + end := reader.position + 8 + if end > len(reader.raw) { + return 0, false + } + value := math.Float64frombits(order.Uint64(reader.raw[reader.position:end])) + reader.position = end + return value, true +} + +func (reader *wkbReader) remaining() int { + return len(reader.raw) - reader.position +} + +func parseWKBType(typeWord uint32) (uint32, wkbDimensions, bool) { + baseType := typeWord & 0x1fffffff + dimensions := wkbDimensions{hasZ: typeWord&0x80000000 != 0, hasM: typeWord&0x40000000 != 0} + hasSRID := typeWord&0x20000000 != 0 + switch { + case baseType >= 3000: + dimensions.hasZ = true + dimensions.hasM = true + baseType -= 3000 + case baseType >= 2000: + dimensions.hasM = true + baseType -= 2000 + case baseType >= 1000: + dimensions.hasZ = true + baseType -= 1000 + } + return baseType, dimensions, hasSRID +} + +func readWKBOrder(reader *wkbReader) (binary.ByteOrder, bool) { + value, ok := reader.readByte() + if !ok { + return nil, false + } + switch value { + case 0: + return binary.BigEndian, true + case 1: + return binary.LittleEndian, true + default: + return nil, false + } +} + +func parseWKBPoints(reader *wkbReader, order binary.ByteOrder, dimensions wkbDimensions) ([][]float64, bool) { + count, ok := reader.readUint32(order) + if !ok { + return nil, false + } + coordinateLength := dimensions.coordinateLength() + required := uint64(count) * uint64(coordinateLength) * 8 + if required > uint64(reader.remaining()) { + return nil, false + } + points := make([][]float64, int(count)) + for pointIndex := range points { + coordinates := make([]float64, coordinateLength) + for coordinateIndex := range coordinates { + value, valueOK := reader.readFloat64(order) + if !valueOK { + return nil, false + } + coordinates[coordinateIndex] = value + } + points[pointIndex] = coordinates + } + return points, true +} + +func readWKBPoint(reader *wkbReader, expected wkbDimensions, depth int) (wkbPoint, bool) { + geometry, _, ok := parseWKBGeometry(reader, depth+1) + if !ok || geometry.kind != 1 || geometry.dimensions != expected { + return wkbPoint{}, false + } + return wkbPoint{coords: geometry.coords, empty: geometry.coords == nil}, true +} + +func parseWKBGeometry(reader *wkbReader, depth int) (wkbGeometry, *uint32, bool) { + if depth > 64 { + return wkbGeometry{}, nil, false + } + order, ok := readWKBOrder(reader) + if !ok { + return wkbGeometry{}, nil, false + } + typeWord, ok := reader.readUint32(order) + if !ok { + return wkbGeometry{}, nil, false + } + baseType, dimensions, hasSRID := parseWKBType(typeWord) + var srid *uint32 + if hasSRID { + value, valueOK := reader.readUint32(order) + if !valueOK { + return wkbGeometry{}, nil, false + } + if value != 0 { + srid = uint32Pointer(value) + } + } + geometry := wkbGeometry{kind: baseType, dimensions: dimensions} + switch baseType { + case 1: + coordinates := make([]float64, dimensions.coordinateLength()) + allNaN := true + for index := range coordinates { + value, valueOK := reader.readFloat64(order) + if !valueOK { + return wkbGeometry{}, nil, false + } + coordinates[index] = value + allNaN = allNaN && math.IsNaN(value) + } + if !allNaN { + geometry.coords = coordinates + } + case 2: + points, pointsOK := parseWKBPoints(reader, order, dimensions) + if !pointsOK { + return wkbGeometry{}, nil, false + } + geometry.points = points + case 3: + count, countOK := reader.readUint32(order) + if !countOK || uint64(count)*4 > uint64(reader.remaining()) { + return wkbGeometry{}, nil, false + } + geometry.rings = make([][][]float64, int(count)) + for index := range geometry.rings { + ring, ringOK := parseWKBPoints(reader, order, dimensions) + if !ringOK { + return wkbGeometry{}, nil, false + } + geometry.rings[index] = ring + } + case 4: + count, countOK := reader.readUint32(order) + if !countOK || uint64(count)*5 > uint64(reader.remaining()) { + return wkbGeometry{}, nil, false + } + geometry.multiPoint = make([]wkbPoint, int(count)) + for index := range geometry.multiPoint { + point, pointOK := readWKBPoint(reader, dimensions, depth) + if !pointOK { + return wkbGeometry{}, nil, false + } + geometry.multiPoint[index] = point + } + case 5: + count, countOK := reader.readUint32(order) + if !countOK || uint64(count)*5 > uint64(reader.remaining()) { + return wkbGeometry{}, nil, false + } + geometry.rings = make([][][]float64, int(count)) + for index := range geometry.rings { + child, _, childOK := parseWKBGeometry(reader, depth+1) + if !childOK || child.kind != 2 { + return wkbGeometry{}, nil, false + } + geometry.rings[index] = child.points + } + case 6: + count, countOK := reader.readUint32(order) + if !countOK || uint64(count)*5 > uint64(reader.remaining()) { + return wkbGeometry{}, nil, false + } + geometry.polygons = make([][][][]float64, int(count)) + for index := range geometry.polygons { + child, _, childOK := parseWKBGeometry(reader, depth+1) + if !childOK || child.kind != 3 { + return wkbGeometry{}, nil, false + } + geometry.polygons[index] = child.rings + } + case 7: + count, countOK := reader.readUint32(order) + if !countOK || uint64(count)*5 > uint64(reader.remaining()) { + return wkbGeometry{}, nil, false + } + geometry.children = make([]wkbGeometry, int(count)) + for index := range geometry.children { + child, _, childOK := parseWKBGeometry(reader, depth+1) + if !childOK { + return wkbGeometry{}, nil, false + } + geometry.children[index] = child + } + default: + return wkbGeometry{}, nil, false + } + return geometry, srid, true +} + +func decodeWKBGeometry(raw []byte) (decodedWKBGeometry, bool) { + reader := &wkbReader{raw: raw} + geometry, srid, ok := parseWKBGeometry(reader, 0) + if !ok || reader.position != len(raw) { + return decodedWKBGeometry{}, false + } + return decodedWKBGeometry{WKT: geometry.wkt(), SRID: srid}, true +} diff --git a/agents/drivers/vastbase-go/spatial_test.go b/agents/drivers/vastbase-go/spatial_test.go new file mode 100644 index 000000000..020ade058 --- /dev/null +++ b/agents/drivers/vastbase-go/spatial_test.go @@ -0,0 +1,79 @@ +package main + +import ( + "encoding/hex" + "reflect" + "testing" +) + +func TestDecodeSpatialValuesMatchesJDBCGeometryShape(t *testing.T) { + point := mustDecodeHex(t, "0101000020E6100000C520B07268195D404E62105839F44340") + tests := []struct { + name string + value any + wkt any + srid *uint32 + }{ + {name: "raw ewkb", value: point, wkt: "POINT(116.397 39.908)", srid: uint32Pointer(4326)}, + {name: "hex ewkb", value: "0x0101000020E6100000C520B07268195D404E62105839F44340", wkt: "POINT(116.397 39.908)", srid: uint32Pointer(4326)}, + {name: "pq text bytes", value: []byte("0101000020E6100000C520B07268195D404E62105839F44340"), wkt: "POINT(116.397 39.908)", srid: uint32Pointer(4326)}, + {name: "ewkt", value: "SRID=3857;POINT(1 2)", wkt: "POINT(1 2)", srid: uint32Pointer(3857)}, + {name: "wkt", value: "POINT(1 2)", wkt: "POINT(1 2)"}, + {name: "null", value: nil, wkt: nil}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + wkt, srid := decodeSpatialValue(test.value) + if !reflect.DeepEqual(wkt, test.wkt) || !reflect.DeepEqual(srid, test.srid) { + t.Fatalf("decodeSpatialValue(%v) = (%v, %v), want (%v, %v)", test.value, wkt, srid, test.wkt, test.srid) + } + }) + } +} + +func TestDecodeSpatialValuesSupportsComplexEWKB(t *testing.T) { + tests := map[string]string{ + "0106000020E610000002000000010300000001000000050000000000000000005D4000000000000044400000000000405D4000000000000044400000000000405D4000000000008044400000000000005D4000000000008044400000000000005D400000000000004440010300000001000000050000000000000000805D4000000000008043400000000000C05D4000000000008043400000000000C05D4000000000000044400000000000805D4000000000000044400000000000805D400000000000804340": "MULTIPOLYGON(((116 40,117 40,117 41,116 41,116 40)),((118 39,119 39,119 40,118 40,118 39)))", + "0107000020E61000000200000001010000000000000000005D4000000000000044400102000000020000000000000000405D4000000000008044400000000000805D400000000000004540": "GEOMETRYCOLLECTION(POINT(116 40),LINESTRING(117 41,118 42))", + } + for encoded, expected := range tests { + decoded, ok := decodeWKBGeometry(mustDecodeHex(t, encoded)) + if !ok || decoded.WKT != expected || decoded.SRID == nil || *decoded.SRID != 4326 { + t.Fatalf("unexpected complex EWKB decode: %+v ok=%v", decoded, ok) + } + } +} + +func TestSpatialDecoderBuildsColumnAndPerCellMetadata(t *testing.T) { + decoder := newSpatialDecoder([]string{"INT4", "public.geometry", "geography(POINT,4326)"}) + row, values, err := decoder.normalizeRow([]any{int64(1), "SRID=4326;POINT(1 2)", nil}) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(row, []any{int64(1), "POINT(1 2)", nil}) { + t.Fatalf("unexpected normalized row: %v", row) + } + if len(values) != 3 || values[0] != nil || values[1] == nil || *values[1] != 4326 || values[2] != nil { + t.Fatalf("unexpected per-cell spatial metadata: %v", values) + } + columns := decoder.columns() + if len(columns) != 2 || columns[0].ColumnIndex != 1 || columns[0].SRID == nil || *columns[0].SRID != 4326 || columns[1].ColumnIndex != 2 || columns[1].SRID != nil { + t.Fatalf("unexpected spatial columns: %+v", columns) + } +} + +func TestMalformedSpatialBytesFallBackToHex(t *testing.T) { + value, srid := decodeSpatialValue(mustDecodeHex(t, "0101000020E6100000C520B072")) + if value != "0x0101000020e6100000c520b072" || srid != nil { + t.Fatalf("unexpected malformed EWKB fallback: value=%v srid=%v", value, srid) + } +} + +func mustDecodeHex(t *testing.T, value string) []byte { + t.Helper() + decoded, err := hex.DecodeString(value) + if err != nil { + t.Fatal(err) + } + return decoded +} diff --git a/agents/drivers/vastbase-go/vastbase_metadata.go b/agents/drivers/vastbase-go/vastbase_metadata.go new file mode 100644 index 000000000..067306c8d --- /dev/null +++ b/agents/drivers/vastbase-go/vastbase_metadata.go @@ -0,0 +1,1234 @@ +package main + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "sort" + "strconv" + "strings" + "time" + + pq "gitcode.com/opengauss/openGauss-connector-go-pq" +) + +const metadataTimeout = 15 * time.Second + +// Escape '_' so only Vastbase internal SYS_/XLOG_ prefixes are hidden; names +// such as SYSTEMS and SYSLOG may be user-created schemas in MySQL mode. +const vastbaseMySQLCompatListSchemasSQL = `SELECT schema_name FROM information_schema.schemata WHERE UPPER(schema_name) <> 'INFORMATION_SCHEMA' AND UPPER(schema_name) NOT LIKE 'SYS\_%' ESCAPE '\' AND UPPER(schema_name) NOT LIKE 'XLOG\_%' ESCAPE '\' ORDER BY schema_name` + +var postgresDataTypes = []string{ + "bigint", "bigserial", "bit", "bit varying", "boolean", "bytea", "char", "character", + "character varying", "date", "decimal", "double precision", "integer", "interval", "json", + "jsonb", "money", "numeric", "real", "smallint", "smallserial", "serial", "text", "time", + "time with time zone", "timestamp", "timestamp with time zone", "uuid", "varchar", "xml", +} + +type vastbaseMode struct { + compatibilityMode string + postgresCatalog bool + mysqlCompat bool + sqlServerIdentity bool +} + +type databaseInfo struct { + Name string `json:"name"` +} + +type tableInfo struct { + Name string `json:"name"` + TableType string `json:"table_type"` + Comment *string `json:"comment"` +} + +type objectInfo struct { + Name string `json:"name"` + ObjectType string `json:"object_type"` + Schema string `json:"schema"` + Comment *string `json:"comment"` + Valid *bool `json:"valid,omitempty"` +} + +type metadataListConstraints struct { + Filter string + Limit int + Offset int + ObjectTypes []string +} + +type columnInfo struct { + Name string `json:"name"` + DataType string `json:"data_type"` + FullDataType string `json:"-"` + IsNullable bool `json:"is_nullable"` + ColumnDefault *string `json:"column_default"` + IsPrimaryKey bool `json:"is_primary_key"` + Extra *string `json:"extra"` + Comment *string `json:"comment"` + NumericPrecision *int `json:"numeric_precision"` + NumericScale *int `json:"numeric_scale"` + CharacterMaximumLength *int `json:"character_maximum_length"` +} + +type indexInfo struct { + Name string `json:"name"` + Columns []string `json:"columns"` + IsUnique bool `json:"is_unique"` + IsPrimary bool `json:"is_primary"` + Filter *string `json:"filter"` + IndexType *string `json:"index_type"` + IncludedColumns []string `json:"included_columns"` + Comment *string `json:"comment"` +} + +func (i indexInfo) MarshalJSON() ([]byte, error) { + type alias indexInfo + value := alias(i) + if value.Columns == nil { + value.Columns = []string{} + } + if value.IncludedColumns == nil { + value.IncludedColumns = []string{} + } + return json.Marshal(value) +} + +type foreignKeyInfo struct { + Name string `json:"name"` + Column string `json:"column"` + RefTable string `json:"ref_table"` + RefColumn string `json:"ref_column"` +} + +type triggerInfo struct { + Name string `json:"name"` + Event string `json:"event"` + Timing string `json:"timing"` +} + +func detectVastbaseMode(db *sql.DB, configuredMySQL bool) vastbaseMode { + if configuredMySQL { + return vastbaseMode{compatibilityMode: "mysql", mysqlCompat: true} + } + mode := vastbaseMode{compatibilityMode: detectDatabaseMode(db)} + mode.postgresCatalog = !catalogExists(db, "sys_catalog.sys_namespace") && catalogExists(db, "pg_catalog.pg_namespace") + if !mode.postgresCatalog { + if mode.compatibilityMode != "" { + mode.mysqlCompat = mode.compatibilityMode == "mysql" + } else { + mode.mysqlCompat = supportsBacktickIdentifiers(db) + } + mode.sqlServerIdentity = !mode.mysqlCompat && catalogExists(db, "sys.identity_columns") + } + return mode +} + +func catalogExists(db *sql.DB, catalog string) bool { + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + rows, err := db.QueryContext(ctx, "SELECT 1 FROM "+catalog+" WHERE 1 = 0") + if err != nil { + return false + } + return rows.Close() == nil +} + +func detectMySQLCompatMode(db *sql.DB) bool { + if databaseMode := detectDatabaseMode(db); databaseMode != "" { + return databaseMode == "mysql" + } + return supportsBacktickIdentifiers(db) +} + +func detectDatabaseMode(db *sql.DB) string { + var databaseMode string + switch err := db.QueryRow("SELECT setting FROM sys_catalog.sys_settings WHERE LOWER(name) = 'database_mode'").Scan(&databaseMode); { + case err == nil: + // Treat database_mode as authoritative when the server exposes it. This + // avoids misclassifying Oracle-compatible servers that also publish + // sql_mode for MySQL syntax toggles such as ANSI_QUOTES. + return strings.ToLower(strings.TrimSpace(databaseMode)) + case errors.Is(err, sql.ErrNoRows): + return "" + default: + return "" + } +} + +func supportsBacktickIdentifiers(db *sql.DB) bool { + var value int + return db.QueryRow("SELECT 1 AS `dbx_identifier_probe`").Scan(&value) == nil +} + +func (s *server) identifierQuote() string { + // Vastbase MySQL compatibility mode follows MySQL identifier quoting; + // other modes retain the PostgreSQL-compatible double quote. + if s.mode.mysqlCompat { + return "`" + } + return `"` +} + +func (s *server) connectionInfo() (map[string]any, error) { + db, err := s.metadataDatabase() + if err != nil { + return nil, err + } + var database, username, version, schema string + err = db.QueryRow("SELECT current_database(), current_user, version(), current_schema()").Scan(&database, &username, &version, &schema) + if err != nil { + return nil, err + } + return map[string]any{ + "database": database, "username": username, "version": version, "schema": schema, + "compatibilityMode": s.mode.compatibilityMode, "mysql_compat_mode": s.mode.mysqlCompat, + "identifierQuote": s.identifierQuote(), + "databaseInfo": map[string]string{ + "productName": "Vastbase", + "productVersion": version, + "unquotedIdentifierCase": "lower", + "quotedIdentifierCase": "mixed", + "driverName": agentDriverName, + "driverVersion": agentDriverVersion, + }, + }, nil +} + +func (s *server) listDatabases() ([]databaseInfo, error) { + queries := []string{ + "SELECT datname FROM sys_catalog.sys_database WHERE NOT datistemplate AND datallowconn ORDER BY datname", + "SELECT datname FROM pg_catalog.pg_database WHERE NOT datistemplate AND datallowconn ORDER BY datname", + "SELECT current_database()", + } + for _, query := range queries { + rows, err := s.metadataQuery(query) + if err != nil { + continue + } + result := []databaseInfo{} + for rows.Next() { + var name string + if rows.Scan(&name) == nil { + result = append(result, databaseInfo{Name: name}) + } + } + err = rows.Err() + _ = rows.Close() + if err == nil && len(result) > 0 { + return result, nil + } + } + return []databaseInfo{{Name: s.params.Database}}, nil +} + +func vastbaseListSchemasSQL(mode vastbaseMode, showSystemSchemas bool) string { + if mode.mysqlCompat { + if showSystemSchemas { + return "SELECT schema_name FROM information_schema.schemata ORDER BY schema_name" + } + return vastbaseMySQLCompatListSchemasSQL + } + if mode.postgresCatalog { + if showSystemSchemas { + return "SELECT nspname FROM pg_catalog.pg_namespace ORDER BY nspname" + } + return "SELECT nspname FROM pg_catalog.pg_namespace WHERE nspname NOT LIKE 'pg_temp_%' AND nspname NOT LIKE 'pg_toast_temp_%' ORDER BY nspname" + } + if showSystemSchemas { + return "SELECT nspname FROM sys_catalog.sys_namespace ORDER BY nspname" + } + return "SELECT nspname FROM sys_catalog.sys_namespace WHERE nspname NOT LIKE 'sys_temp_%' AND nspname NOT LIKE 'sys_toast_temp_%' ORDER BY nspname" +} + +func (s *server) listSchemas(visible []string, showSystemSchemas bool) ([]string, error) { + query := vastbaseListSchemasSQL(s.mode, showSystemSchemas) + rows, err := s.metadataQuery(query) + if err != nil { + return nil, err + } + defer rows.Close() + allowed := stringSet(visible) + result := []string{} + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, err + } + if len(allowed) == 0 || allowed[strings.ToLower(name)] { + result = append(result, name) + } + } + return result, rows.Err() +} + +func (s *server) listTables(schema string, constraints metadataListConstraints) ([]tableInfo, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + if !constraintsAllowsTableLike(constraints) { + return []tableInfo{}, nil + } + catalog := "sys_catalog" + if s.mode.postgresCatalog { + catalog = "pg_catalog" + } + query := fmt.Sprintf(`SELECT c.relname, +CASE c.relkind WHEN 'r' THEN 'TABLE' WHEN 'p' THEN 'TABLE' WHEN 'v' THEN 'VIEW' WHEN 'm' THEN 'MATERIALIZED_VIEW' WHEN 'f' THEN 'FOREIGN_TABLE' ELSE 'TABLE' END, +obj_description(c.oid) +FROM %s.%s_class c +JOIN %s.%s_namespace n ON n.oid = c.relnamespace +WHERE n.nspname = $1 AND c.relkind IN ('r','p','v','m','f') ORDER BY c.relname`, catalog, catalogPrefix(catalog), catalog, catalogPrefix(catalog)) + rows, err := s.cachedListTablesQuery(query, effective) + if err != nil { + return nil, err + } + defer rows.Close() + result := []tableInfo{} + for rows.Next() { + var name, kind string + var comment sql.NullString + if err := rows.Scan(&name, &kind, &comment); err != nil { + return nil, err + } + item := tableInfo{Name: name, TableType: normalizeTableType(kind), Comment: nullStringPtr(comment)} + if constraintsMatch(constraints, item.Name, item.TableType) { + result = append(result, item) + } + } + return pageTables(result, constraints), rows.Err() +} + +func (s *server) cachedListTablesQuery(query, schema string) (*sql.Rows, error) { + if s.connectionRuntime != nil && !s.sessionAffinity { + return s.connectionRuntime.queryListTables(query, schema) + } + db, err := s.requireDB() + if err != nil { + return nil, err + } + if s.listTablesStatement == nil { + statement, prepareErr := db.Prepare(query) + if prepareErr != nil { + return nil, prepareErr + } + s.listTablesStatement = statement + } + return s.listTablesStatement.Query(schema) +} + +func (s *server) getTableComment(schema, table string) (*string, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + catalog := "sys_catalog" + if s.mode.postgresCatalog { + catalog = "pg_catalog" + } + prefix := catalogPrefix(catalog) + query := fmt.Sprintf(`SELECT obj_description(c.oid) +FROM %s.%s_class c +JOIN %s.%s_namespace n ON n.oid = c.relnamespace +WHERE n.nspname = %s AND c.relname = %s AND c.relkind IN ('r','p','v','m','f') +LIMIT 1`, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(table)) + var comment sql.NullString + if err := s.requireDBQueryRow(query, &comment); err != nil { + if err == sql.ErrNoRows { + return nil, nil + } + return nil, err + } + return nullStringPtr(comment), nil +} + +func (s *server) listObjects(schema string, constraints metadataListConstraints) ([]objectInfo, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + tables, err := s.listTables(effective, metadataListConstraints{}) + if err != nil { + return nil, err + } + result := make([]objectInfo, 0, len(tables)) + for _, table := range tables { + result = append(result, objectInfo{Name: table.Name, ObjectType: table.TableType, Schema: effective, Comment: table.Comment}) + } + if !s.mode.mysqlCompat { + catalog := "sys_catalog" + function := "sys" + if s.mode.postgresCatalog { + catalog, function = "pg_catalog", "pg" + } + query := fmt.Sprintf(`SELECT p.proname, CASE WHEN p.prorettype = 2278 THEN 'PROCEDURE' ELSE 'FUNCTION' END, d.description +FROM %s.%s_proc p JOIN %s.%s_namespace n ON n.oid = p.pronamespace +LEFT JOIN %s.%s_description d ON d.objoid = p.oid AND d.objsubid = 0 +WHERE n.nspname = %s ORDER BY p.proname`, catalog, function, catalog, function, catalog, function, quoteLiteral(effective)) + rows, queryErr := s.metadataQuery(query) + if queryErr == nil { + for rows.Next() { + var name, kind string + var comment sql.NullString + if rows.Scan(&name, &kind, &comment) == nil { + result = append(result, objectInfo{Name: name, ObjectType: kind, Schema: effective, Comment: nullStringPtr(comment)}) + } + } + _ = rows.Close() + } + } + filtered := result[:0] + for _, item := range result { + if constraintsMatch(constraints, item.Name, item.ObjectType) { + filtered = append(filtered, item) + } + } + sort.SliceStable(filtered, func(i, j int) bool { + if objectOrder(filtered[i].ObjectType) != objectOrder(filtered[j].ObjectType) { + return objectOrder(filtered[i].ObjectType) < objectOrder(filtered[j].ObjectType) + } + return filtered[i].Name < filtered[j].Name + }) + return pageObjects(filtered, constraints), nil +} + +func (s *server) completionAssistantSearch(request completionAssistantRequest) (completionAssistantResponse, error) { + limit := request.MaxResults + if limit <= 0 || limit > 1000 { + limit = 100 + } + kinds := stringSet(request.ObjectKinds) + candidates := make([]completionAssistantCandidate, 0, limit+1) + if kinds["column"] && request.ParentName != "" { + schema := request.ParentSchema + if schema == "" { + schema = request.Schema + } + columns, err := s.getColumns(schema, request.ParentName) + if err != nil { + return completionAssistantResponse{}, err + } + for _, column := range columns { + if !completionNameMatches(column.Name, request) { + continue + } + dataType := column.DataType + candidates = append(candidates, completionAssistantCandidate{ + Name: column.Name, Kind: "COLUMN", Schema: stringPtr(schema), ParentSchema: stringPtr(schema), + ParentName: stringPtr(request.ParentName), Comment: column.Comment, DataType: &dataType, + }) + } + } else { + schemas := []string{request.Schema} + if request.GlobalSearch { + visible, err := s.listSchemas(nil, false) + if err != nil { + return completionAssistantResponse{}, err + } + schemas = visible + } + objectTypes := request.ObjectKinds + for _, schema := range schemas { + objects, err := s.listObjects(schema, metadataListConstraints{ObjectTypes: objectTypes}) + if err != nil { + return completionAssistantResponse{}, err + } + for _, object := range objects { + if !completionNameMatches(object.Name, request) { + continue + } + candidates = append(candidates, completionAssistantCandidate{Name: object.Name, Kind: object.ObjectType, Schema: stringPtr(object.Schema), Comment: object.Comment}) + if len(candidates) > limit { + return completionAssistantResponse{Candidates: candidates[:limit], Incomplete: true}, nil + } + } + } + } + incomplete := len(candidates) > limit + if incomplete { + candidates = candidates[:limit] + } + if candidates == nil { + candidates = []completionAssistantCandidate{} + } + return completionAssistantResponse{Candidates: candidates, Incomplete: incomplete}, nil +} + +func completionNameMatches(name string, request completionAssistantRequest) bool { + mask := request.Mask + if mask == "" { + return true + } + if !request.CaseSensitive { + name = strings.ToLower(name) + mask = strings.ToLower(mask) + } + if strings.EqualFold(request.MatchMode, "contains") { + return strings.Contains(name, mask) + } + return strings.HasPrefix(name, mask) +} + +func (s *server) getColumns(schema, table string) ([]columnInfo, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + primary, _ := s.primaryKeys(effective, table) + if s.mode.mysqlCompat { + return s.informationSchemaColumns(effective, table, primary) + } + catalog, prefix := "sys_catalog", "sys" + if s.mode.postgresCatalog { + catalog, prefix = "pg_catalog", "pg" + return s.queryCatalogColumns(effective, table, primary, catalog, prefix, "pg_get_expr") + } + expression := "sys_get_expr" + if s.usePgDefaultExpression { + expression = "pg_get_expr" + } + result, err := s.queryCatalogColumns(effective, table, primary, catalog, prefix, expression) + if err != nil && expression == "sys_get_expr" && isUndefinedFunction(err, expression) { + // Some V8R6 PostgreSQL-mode databases keep sys_catalog while adbin is + // pg_node_tree. Cache the compatible function after the exact failure. + s.usePgDefaultExpression = true + return s.queryCatalogColumns(effective, table, primary, catalog, prefix, "pg_get_expr") + } + return result, err +} + +func (s *server) queryCatalogColumns( + schema, table string, + primary map[string]bool, + catalog, prefix, expression string, +) ([]columnInfo, error) { + identityExpression := "a.attidentity" + if s.catalogIdentityUnsupported { + identityExpression = "CAST(NULL AS varchar(1)) AS attidentity" + } + query := fmt.Sprintf(`SELECT a.attname, format_type(a.atttypid, a.atttypmod), NOT a.attnotnull, + %s(ad.adbin, ad.adrelid), col_description(a.attrelid, a.attnum), + CASE WHEN t.typname = 'numeric' AND a.atttypmod > 0 THEN ((a.atttypmod - 4) >> 16) & 65535 END, + CASE WHEN t.typname = 'numeric' AND a.atttypmod > 0 THEN (a.atttypmod - 4) & 65535 END, + CASE WHEN t.typname IN ('varchar','bpchar') AND a.atttypmod > 0 THEN a.atttypmod - 4 END, + %s + FROM %s.%s_attribute a JOIN %s.%s_type t ON t.oid = a.atttypid + JOIN %s.%s_class c ON c.oid = a.attrelid JOIN %s.%s_namespace n ON n.oid = c.relnamespace + LEFT JOIN %s.%s_attrdef ad ON ad.adrelid = a.attrelid AND ad.adnum = a.attnum +WHERE n.nspname = %s AND c.relname = %s AND a.attnum > 0 AND NOT a.attisdropped ORDER BY a.attnum`, expression, identityExpression, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, quoteLiteral(schema), quoteLiteral(table)) + rows, err := s.metadataQuery(query) + if err != nil && !s.catalogIdentityUnsupported && isUndefinedColumn(err, "attidentity") { + s.catalogIdentityUnsupported = true + return s.queryCatalogColumns(schema, table, primary, catalog, prefix, expression) + } + if err != nil { + return nil, err + } + defer rows.Close() + result := []columnInfo{} + for rows.Next() { + var name, dataType string + var nullable bool + var defaultValue, comment, identity sql.NullString + var precision, scale, length sql.NullInt64 + if err := rows.Scan(&name, &dataType, &nullable, &defaultValue, &comment, &precision, &scale, &length, &identity); err != nil { + return nil, err + } + result = append(result, columnInfo{Name: name, DataType: dataType, IsNullable: nullable, ColumnDefault: nullStringPtr(defaultValue), IsPrimaryKey: primary[strings.ToLower(name)], Extra: vastbaseIdentityClause(identity.String), Comment: nullStringPtr(comment), NumericPrecision: nullIntPtr(precision), NumericScale: nullIntPtr(scale), CharacterMaximumLength: nullIntPtr(length)}) + } + if err := rows.Err(); err != nil { + return nil, err + } + if s.mode.sqlServerIdentity { + s.applyIdentityMetadata(schema, table, result) + } + return result, nil +} + +func isUndefinedFunction(err error, functionName string) bool { + var driverError *pq.Error + undefined := errors.As(err, &driverError) && string(driverError.Code) == "42883" + normalized := strings.ToLower(err.Error()) + undefined = undefined || strings.Contains(normalized, "does not exist") || strings.Contains(normalized, "不存在") + return undefined && strings.Contains(normalized, strings.ToLower(functionName)) +} + +func isUndefinedColumn(err error, columnName string) bool { + var driverError *pq.Error + undefined := errors.As(err, &driverError) && string(driverError.Code) == "42703" + normalized := strings.ToLower(err.Error()) + undefined = undefined || strings.Contains(normalized, "does not exist") || strings.Contains(normalized, "不存在") + return undefined && strings.Contains(normalized, strings.ToLower(columnName)) +} + +func (s *server) informationSchemaColumns(schema, table string, primary map[string]bool) ([]columnInfo, error) { + includeColumnType := !s.infoColumnTypeUnsupported + includeUdtName := !s.infoUdtNameUnsupported + for { + result, err := s.queryInformationSchemaColumns(schema, table, primary, includeColumnType, includeUdtName) + if err == nil { + return result, nil + } + switch { + case includeColumnType && isUndefinedColumn(err, "column_type"): + includeColumnType = false + s.infoColumnTypeUnsupported = true + case includeUdtName && isUndefinedColumn(err, "udt_name"): + includeUdtName = false + s.infoUdtNameUnsupported = true + default: + return nil, err + } + } +} + +func (s *server) queryInformationSchemaColumns(schema, table string, primary map[string]bool, includeColumnType, includeUdtName bool) ([]columnInfo, error) { + var fullDataTypeExpression string + switch { + case includeColumnType && includeUdtName: + fullDataTypeExpression = `CASE + WHEN UPPER(TRIM(c.data_type)) IN ('USER-DEFINED', 'USER_DEFINED') + AND UPPER(COALESCE(NULLIF(TRIM(c.column_type), ''), 'USER-DEFINED')) IN ('USER-DEFINED', 'USER_DEFINED') + THEN c.udt_name + ELSE c.column_type + END` + case includeColumnType: + fullDataTypeExpression = "c.column_type" + case includeUdtName: + fullDataTypeExpression = `CASE + WHEN UPPER(TRIM(c.data_type)) IN ('USER-DEFINED', 'USER_DEFINED') THEN c.udt_name + END AS column_type` + default: + fullDataTypeExpression = "NULL AS column_type" + } + query := fmt.Sprintf(`SELECT c.column_name, c.data_type, %s, c.is_nullable, c.column_default, + col_description(a.attrelid, a.attnum), c.numeric_precision, c.numeric_scale, c.character_maximum_length + FROM information_schema.columns c + LEFT JOIN sys_catalog.sys_namespace n ON n.nspname = c.table_schema + LEFT JOIN sys_catalog.sys_class rel ON rel.relnamespace = n.oid AND rel.relname = c.table_name + LEFT JOIN sys_catalog.sys_attribute a ON a.attrelid = rel.oid AND a.attname = c.column_name AND a.attnum > 0 AND NOT a.attisdropped + WHERE c.table_schema = %s AND c.table_name = %s ORDER BY c.ordinal_position`, fullDataTypeExpression, quoteLiteral(schema), quoteLiteral(table)) + rows, err := s.metadataQuery(query) + if err != nil { + return nil, err + } + defer rows.Close() + result := []columnInfo{} + for rows.Next() { + var name, dataType, nullable string + var fullDataType, defaultValue, comment sql.NullString + var precision, scale, length sql.NullInt64 + if err := rows.Scan(&name, &dataType, &fullDataType, &nullable, &defaultValue, &comment, &precision, &scale, &length); err != nil { + return nil, err + } + if parsed := boundedVarcharLength(dataType); parsed != nil && !length.Valid { + length = sql.NullInt64{Int64: int64(*parsed), Valid: true} + } + result = append(result, columnInfo{Name: name, DataType: dataType, FullDataType: fullDataType.String, IsNullable: strings.EqualFold(nullable, "YES"), ColumnDefault: nullStringPtr(defaultValue), IsPrimaryKey: primary[strings.ToLower(name)], Comment: nullStringPtr(comment), NumericPrecision: nullIntPtr(precision), NumericScale: nullIntPtr(scale), CharacterMaximumLength: nullIntPtr(length)}) + } + return result, rows.Err() +} + +func (s *server) listIndexes(schema, table string) ([]indexInfo, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + catalog, prefix := "sys_catalog", "sys" + if s.mode.postgresCatalog { + catalog, prefix = "pg_catalog", "pg" + } + query := vastbaseListIndexesQuery(catalog, prefix, effective, table) + rows, err := s.metadataQuery(query) + if err != nil { + return nil, err + } + defer rows.Close() + byName := map[string]*indexInfo{} + order := []string{} + for rows.Next() { + var name, kind, column string + var unique, primary bool + var ordinal int + if err := rows.Scan(&name, &kind, &unique, &primary, &column, &ordinal); err != nil { + return nil, err + } + item := byName[name] + if item == nil { + item = &indexInfo{Name: name, IsUnique: unique, IsPrimary: primary, IndexType: stringPtr(kind), Columns: []string{}, IncludedColumns: []string{}} + byName[name] = item + order = append(order, name) + } + item.Columns = append(item.Columns, column) + } + result := make([]indexInfo, 0, len(order)) + for _, name := range order { + result = append(result, *byName[name]) + } + return result, rows.Err() +} + +func vastbaseListIndexesQuery(catalog, prefix, schema, table string) string { + return fmt.Sprintf(`SELECT i.relname, am.amname, ix.indisunique, ix.indisprimary, a.attname, pos.n +FROM %s.%s_index ix JOIN %s.%s_class t ON t.oid = ix.indrelid +JOIN %s.%s_class i ON i.oid = ix.indexrelid JOIN %s.%s_namespace n ON n.oid = t.relnamespace +JOIN %s.%s_am am ON am.oid = i.relam +JOIN unnest(ix.indkey) WITH ORDINALITY AS pos(attnum,n) ON true +JOIN %s.%s_attribute a ON a.attrelid = t.oid AND a.attnum = pos.attnum +WHERE n.nspname = %s AND t.relname = %s ORDER BY i.relname, pos.n`, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, quoteLiteral(schema), quoteLiteral(table)) +} + +func vastbaseCatalogFunction(catalog, sysFunction, postgresFunction string) string { + // Vastbase compatibility modes expose different catalog-qualified deparser names. + if catalog == "pg_catalog" { + return "pg_catalog." + postgresFunction + } + return "sys_catalog." + sysFunction +} + +func (s *server) listForeignKeys(schema, table string) ([]foreignKeyInfo, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + query := `SELECT fk.constraint_name, fk.column_name, pk.table_name, pk.column_name +FROM information_schema.table_constraints tc +JOIN information_schema.key_column_usage fk ON fk.constraint_schema = tc.constraint_schema AND fk.constraint_name = tc.constraint_name AND fk.table_schema = tc.table_schema AND fk.table_name = tc.table_name +JOIN information_schema.referential_constraints rc ON rc.constraint_schema = tc.constraint_schema AND rc.constraint_name = tc.constraint_name +JOIN information_schema.key_column_usage pk ON pk.constraint_schema = rc.unique_constraint_schema AND pk.constraint_name = rc.unique_constraint_name AND pk.ordinal_position = fk.position_in_unique_constraint +WHERE tc.table_schema = ` + quoteLiteral(effective) + ` AND tc.table_name = ` + quoteLiteral(table) + ` AND tc.constraint_type = 'FOREIGN KEY' ORDER BY fk.constraint_name, fk.ordinal_position` + rows, err := s.metadataQuery(query) + if err != nil { + return nil, err + } + defer rows.Close() + result := []foreignKeyInfo{} + for rows.Next() { + var item foreignKeyInfo + if err := rows.Scan(&item.Name, &item.Column, &item.RefTable, &item.RefColumn); err != nil { + return nil, err + } + result = append(result, item) + } + return result, rows.Err() +} + +func (s *server) listTriggers(schema, table string) ([]triggerInfo, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + catalog, prefix := "sys_catalog", "sys" + if s.mode.postgresCatalog { + catalog, prefix = "pg_catalog", "pg" + } + query := fmt.Sprintf(`SELECT tg.tgname, +trim(trailing ',' FROM (CASE WHEN (tg.tgtype & 4) <> 0 THEN 'INSERT,' ELSE '' END || CASE WHEN (tg.tgtype & 8) <> 0 THEN 'DELETE,' ELSE '' END || CASE WHEN (tg.tgtype & 16) <> 0 THEN 'UPDATE,' ELSE '' END || CASE WHEN (tg.tgtype & 32) <> 0 THEN 'TRUNCATE,' ELSE '' END)), tg.tgtype +FROM %s.%s_trigger tg JOIN %s.%s_class c ON c.oid = tg.tgrelid JOIN %s.%s_namespace n ON n.oid = c.relnamespace +WHERE n.nspname = %s AND c.relname = %s AND NOT tg.tgisinternal ORDER BY tg.tgname`, catalog, prefix, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(table)) + rows, err := s.metadataQuery(query) + if err != nil { + return nil, err + } + defer rows.Close() + result := []triggerInfo{} + for rows.Next() { + var name, event string + var triggerType int + if err := rows.Scan(&name, &event, &triggerType); err != nil { + return nil, err + } + result = append(result, triggerInfo{Name: name, Event: event, Timing: decodeTriggerTiming(triggerType)}) + } + return result, rows.Err() +} + +func (s *server) getObjectSource(schema, name, objectType string) (map[string]any, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + source := "" + kind := strings.ToUpper(objectType) + if kind == "VIEW" || kind == "MATERIALIZED_VIEW" { + if s.mode.mysqlCompat { + err = s.requireDBQueryRow("SELECT view_definition FROM information_schema.views WHERE table_schema = "+quoteLiteral(effective)+" AND table_name = "+quoteLiteral(name), &source) + } else { + catalog, prefix, function := "sys_catalog", "sys", "sys_get_viewdef" + if s.mode.postgresCatalog { + catalog, prefix, function = "pg_catalog", "pg", "pg_get_viewdef" + } + query := fmt.Sprintf("SELECT %s(c.oid) FROM %s.%s_class c JOIN %s.%s_namespace n ON n.oid=c.relnamespace WHERE n.nspname=%s AND c.relname=%s LIMIT 1", function, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(name)) + err = s.requireDBQueryRow(query, &source) + } + } else if kind == "FUNCTION" || kind == "PROCEDURE" { + catalog, prefix, function := "sys_catalog", "sys", "sys_get_functiondef" + if s.mode.postgresCatalog { + catalog, prefix, function = "pg_catalog", "pg", "pg_get_functiondef" + } + query := fmt.Sprintf("SELECT %s(p.oid) FROM %s.%s_proc p JOIN %s.%s_namespace n ON n.oid=p.pronamespace WHERE n.nspname=%s AND p.proname=%s ORDER BY CASE WHEN p.prorettype=2278 THEN 0 ELSE 1 END LIMIT 1", function, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(name)) + err = s.requireDBQueryRow(query, &source) + } + if err != nil && err != sql.ErrNoRows { + return nil, err + } + source = normalizeAgentObjectSource(source) + return map[string]any{"name": name, "object_type": objectType, "schema": effective, "source": source}, nil +} + +func (s *server) getTableDDL(schema, table string) (string, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return "", err + } + columns, err := s.getColumns(effective, table) + if err != nil { + return "", err + } + tableComment, _ := s.getTableComment(effective, table) + ddl := renderTableDDL(effective, table, columns, tableComment) + ddl, err = s.appendTableIndexDDL(effective, table, ddl) + if err != nil { + return "", err + } + ddl, err = s.appendTableTriggerDDL(effective, table, ddl) + if err != nil { + return "", err + } + return ddl, nil +} + +func renderTableDDL(schema, table string, columns []columnInfo, tableComment *string) string { + definitions := make([]string, 0, len(columns)+1) + primary := []string{} + for _, column := range columns { + definitions = append(definitions, columnDDLDefinition(column)) + if column.IsPrimaryKey { + primary = append(primary, quoteIdentifier(column.Name)) + } + } + if len(primary) > 0 { + definitions = append(definitions, "PRIMARY KEY ("+strings.Join(primary, ", ")+")") + } + qualifiedTable := quoteIdentifier(schema) + "." + quoteIdentifier(table) + ddl := "CREATE TABLE " + qualifiedTable + " (\n " + strings.Join(definitions, ",\n ") + "\n);" + if tableComment != nil && strings.TrimSpace(*tableComment) != "" { + ddl += "\nCOMMENT ON TABLE " + qualifiedTable + " IS " + quoteLiteral(*tableComment) + ";" + } + for _, column := range columns { + if column.Comment == nil || strings.TrimSpace(*column.Comment) == "" { + continue + } + ddl += "\nCOMMENT ON COLUMN " + qualifiedTable + "." + quoteIdentifier(column.Name) + " IS " + quoteLiteral(*column.Comment) + ";" + } + return ddl +} + +func (s *server) appendTableIndexDDL(schema, table, ddl string) (string, error) { + definitions, err := s.listIndexDefinitions(schema, table) + if err != nil { + return "", err + } + return appendDDLStatements(ddl, definitions), nil +} + +func (s *server) appendTableTriggerDDL(schema, table, ddl string) (string, error) { + definitions, err := s.listTriggerDefinitions(schema, table) + if err != nil { + return "", err + } + return appendDDLStatements(ddl, definitions), nil +} + +func (s *server) listIndexDefinitions(schema, table string) ([]string, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + catalog, prefix := "sys_catalog", "sys" + if s.mode.postgresCatalog { + catalog, prefix = "pg_catalog", "pg" + } + indexDefinitionFunction := vastbaseCatalogFunction(catalog, "sys_get_indexdef", "pg_get_indexdef") + query := fmt.Sprintf(`SELECT i.relname, %s(ix.indexrelid, 0, true), obj_description(i.oid) +FROM %s.%s_index ix JOIN %s.%s_class t ON t.oid = ix.indrelid +JOIN %s.%s_class i ON i.oid = ix.indexrelid JOIN %s.%s_namespace n ON n.oid = t.relnamespace +WHERE n.nspname = %s AND t.relname = %s AND NOT ix.indisprimary ORDER BY i.relname`, indexDefinitionFunction, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(table)) + rows, err := s.metadataQuery(query) + if err != nil { + return nil, err + } + defer rows.Close() + result := []string{} + for rows.Next() { + var name string + var definition string + var comment sql.NullString + if err := rows.Scan(&name, &definition, &comment); err != nil { + return nil, err + } + if strings.TrimSpace(definition) != "" { + result = append(result, definition) + } + if comment.Valid && strings.TrimSpace(comment.String) != "" { + result = append(result, "COMMENT ON INDEX "+quoteIdentifier(effective)+"."+quoteIdentifier(name)+" IS "+quoteLiteral(comment.String)) + } + } + return result, rows.Err() +} + +func (s *server) listTriggerDefinitions(schema, table string) ([]string, error) { + effective, err := s.effectiveSchema(schema) + if err != nil { + return nil, err + } + catalog, prefix := "sys_catalog", "sys" + if s.mode.postgresCatalog { + catalog, prefix = "pg_catalog", "pg" + } + triggerDefinitionFunction := vastbaseCatalogFunction(catalog, "sys_get_triggerdef", "pg_get_triggerdef") + query := fmt.Sprintf(`SELECT %s(tg.oid, true) +FROM %s.%s_trigger tg JOIN %s.%s_class c ON c.oid = tg.tgrelid JOIN %s.%s_namespace n ON n.oid = c.relnamespace +WHERE n.nspname = %s AND c.relname = %s AND NOT tg.tgisinternal ORDER BY tg.tgname`, triggerDefinitionFunction, catalog, prefix, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(table)) + rows, err := s.metadataQuery(query) + if err != nil { + return nil, err + } + defer rows.Close() + result := []string{} + for rows.Next() { + var definition string + if err := rows.Scan(&definition); err != nil { + return nil, err + } + if strings.TrimSpace(definition) != "" { + result = append(result, definition) + } + } + return result, rows.Err() +} + +func appendDDLStatements(ddl string, statements []string) string { + for _, statement := range statements { + ddl = appendDDLStatement(ddl, statement) + } + return ddl +} + +func appendDDLStatement(ddl, statement string) string { + ddl = strings.TrimRight(ddl, "\r\n\t ") + statement = ensureStatementTerminator(statement) + if statement == "" { + return ddl + } + if ddl == "" { + return statement + } + if !strings.HasSuffix(ddl, ";") { + ddl += ";" + } + return ddl + "\n\n" + statement +} + +func ensureStatementTerminator(statement string) string { + trimmed := strings.TrimSpace(statement) + if trimmed == "" || strings.HasSuffix(trimmed, ";") { + return trimmed + } + return trimmed + ";" +} + +func columnDDLDefinition(column columnInfo) string { + definition := quoteIdentifier(column.Name) + " " + columnDDLDataType(column) + if column.Extra != nil && *column.Extra != "" { + // Identity clauses belong immediately after the data type in both + // PostgreSQL-compatible and SQL Server-compatible Vastbase modes. + definition += " " + *column.Extra + } + if !column.IsNullable { + definition += " NOT NULL" + } + if column.ColumnDefault != nil && *column.ColumnDefault != "" { + definition += " DEFAULT " + *column.ColumnDefault + } + return definition +} + +func columnDDLDataType(column columnInfo) string { + if fullDataType := strings.TrimSpace(column.FullDataType); fullDataType != "" { + return fullDataType + } + dataType := strings.TrimSpace(column.DataType) + if strings.Contains(dataType, "(") { + return dataType + } + normalized := strings.Join(strings.Fields(strings.ToLower(dataType)), " ") + switch normalized { + case "varchar", "character varying", "char", "character": + if column.CharacterMaximumLength != nil && *column.CharacterMaximumLength > 0 { + return fmt.Sprintf("%s(%d)", dataType, *column.CharacterMaximumLength) + } + case "numeric", "decimal": + if column.NumericPrecision != nil && *column.NumericPrecision > 0 { + if column.NumericScale != nil { + return fmt.Sprintf("%s(%d,%d)", dataType, *column.NumericPrecision, *column.NumericScale) + } + return fmt.Sprintf("%s(%d)", dataType, *column.NumericPrecision) + } + } + return dataType +} + +func vastbaseIdentityClause(code string) *string { + var clause string + switch strings.ToLower(strings.TrimSpace(code)) { + case "a": + clause = "GENERATED ALWAYS AS IDENTITY" + case "d": + clause = "GENERATED BY DEFAULT AS IDENTITY" + default: + return nil + } + return &clause +} + +func (s *server) getExplainInfo(sqlText string) (string, error) { + rows, err := s.metadataQuery("EXPLAIN " + trimStatementSQL(sqlText)) + if err != nil { + return "", err + } + defer rows.Close() + lines := []string{} + for rows.Next() { + var line string + if err := rows.Scan(&line); err != nil { + return "", err + } + lines = append(lines, line) + } + return strings.Join(lines, "\n"), rows.Err() +} + +func (s *server) metadataQuery(query string) (*sql.Rows, error) { + db, err := s.metadataDatabase() + if err != nil { + return nil, err + } + // These are bounded, internally generated statements. Calling Query without + // arguments keeps pq on its single-round-trip simple-query path. + return db.Query(query) +} + +func (s *server) requireDBQueryRow(query string, destination ...any) error { + db, err := s.metadataDatabase() + if err != nil { + return err + } + ctx, cancel := context.WithTimeout(context.Background(), metadataTimeout) + defer cancel() + return db.QueryRowContext(ctx, query).Scan(destination...) +} + +func (s *server) effectiveSchema(schema string) (string, error) { + if strings.TrimSpace(schema) != "" { + return strings.TrimSpace(schema), nil + } + var current sql.NullString + if err := s.requireDBQueryRow("SELECT current_schema()", ¤t); err == nil && current.Valid && current.String != "" { + return current.String, nil + } + if s.params.Username != "" { + return s.params.Username, nil + } + return "public", nil +} + +func (s *server) primaryKeys(schema, table string) (map[string]bool, error) { + query := `SELECT kcu.column_name FROM information_schema.table_constraints tc +JOIN information_schema.key_column_usage kcu ON kcu.constraint_schema=tc.constraint_schema AND kcu.constraint_name=tc.constraint_name AND kcu.table_schema=tc.table_schema AND kcu.table_name=tc.table_name +WHERE tc.table_schema=` + quoteLiteral(schema) + ` AND tc.table_name=` + quoteLiteral(table) + ` AND tc.constraint_type='PRIMARY KEY' ORDER BY kcu.ordinal_position` + rows, err := s.metadataQuery(query) + if err != nil { + return nil, err + } + defer rows.Close() + result := map[string]bool{} + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, err + } + result[strings.ToLower(name)] = true + } + return result, rows.Err() +} + +func (s *server) applyIdentityMetadata(schema, table string, columns []columnInfo) { + query := `SELECT a.attname, ic.seed_value, ic.increment_value FROM sys.identity_columns ic +JOIN sys_catalog.sys_class c ON c.oid=ic.object_id JOIN sys_catalog.sys_namespace n ON n.oid=c.relnamespace +JOIN sys_catalog.sys_attribute a ON a.attrelid=c.oid AND a.attnum=ic.column_id +WHERE n.nspname=` + quoteLiteral(schema) + ` AND c.relname=` + quoteLiteral(table) + rows, err := s.metadataQuery(query) + if err != nil { + s.mode.sqlServerIdentity = false + return + } + defer rows.Close() + byName := map[string]*columnInfo{} + for i := range columns { + byName[strings.ToLower(columns[i].Name)] = &columns[i] + } + for rows.Next() { + var name string + var seed, increment sql.NullString + if rows.Scan(&name, &seed, &increment) == nil { + if column := byName[strings.ToLower(name)]; column != nil { + extra := "IDENTITY" + if seed.Valid && increment.Valid { + extra = "IDENTITY(" + seed.String + "," + increment.String + ")" + } + column.Extra = &extra + } + } + } +} + +func catalogPrefix(catalog string) string { + if catalog == "pg_catalog" { + return "pg" + } + return "sys" +} + +func normalizeTableType(value string) string { + normalized := strings.ToUpper(strings.ReplaceAll(strings.TrimSpace(value), " ", "_")) + switch normalized { + case "BASE_TABLE", "PARTITIONED_TABLE": + return "TABLE" + case "MATERIALIZED_VIEW", "FOREIGN_TABLE", "VIEW", "TABLE": + return normalized + default: + return "TABLE" + } +} + +func decodeTriggerTiming(triggerType int) string { + if triggerType&(1<<6) != 0 { + return "INSTEAD OF" + } + if triggerType&(1<<1) != 0 { + return "BEFORE" + } + return "AFTER" +} + +func boundedVarcharLength(dataType string) *int { + lower := strings.ToLower(strings.TrimSpace(dataType)) + for _, prefix := range []string{"varchar", "character varying"} { + if strings.HasPrefix(lower, prefix) { + value := strings.TrimSpace(strings.TrimSuffix(strings.TrimPrefix(lower, prefix), ")")) + value = strings.TrimPrefix(value, "(") + if number, err := strconv.Atoi(strings.TrimSpace(value)); err == nil && number >= 0 { + return &number + } + } + } + return nil +} + +func constraintsAllowsTableLike(constraints metadataListConstraints) bool { + if len(constraints.ObjectTypes) == 0 { + return true + } + for _, kind := range constraints.ObjectTypes { + switch normalizeTableType(kind) { + case "TABLE", "VIEW", "MATERIALIZED_VIEW", "FOREIGN_TABLE": + return true + } + } + return false +} + +func constraintsMatch(constraints metadataListConstraints, name, kind string) bool { + if filter := strings.TrimSpace(constraints.Filter); filter != "" && !strings.Contains(strings.ToLower(name), strings.ToLower(filter)) { + return false + } + if len(constraints.ObjectTypes) == 0 { + return true + } + for _, allowed := range constraints.ObjectTypes { + if strings.EqualFold(normalizeTableType(allowed), normalizeTableType(kind)) || strings.EqualFold(allowed, kind) { + return true + } + } + return false +} + +func pageTables(items []tableInfo, constraints metadataListConstraints) []tableInfo { + start, end := pageBounds(len(items), constraints.Offset, constraints.Limit) + return items[start:end] +} + +func pageObjects(items []objectInfo, constraints metadataListConstraints) []objectInfo { + start, end := pageBounds(len(items), constraints.Offset, constraints.Limit) + return items[start:end] +} + +func pageBounds(length, offset, limit int) (int, int) { + if offset < 0 { + offset = 0 + } + if offset > length { + offset = length + } + end := length + if limit > 0 && offset+limit < end { + end = offset + limit + } + return offset, end +} + +func objectOrder(kind string) int { + switch strings.ToUpper(kind) { + case "TABLE": + return 0 + case "VIEW": + return 1 + case "MATERIALIZED_VIEW": + return 2 + case "FOREIGN_TABLE": + return 3 + case "PROCEDURE": + return 4 + case "FUNCTION": + return 5 + default: + return 9 + } +} + +func stringSet(values []string) map[string]bool { + result := map[string]bool{} + for _, value := range values { + result[strings.ToLower(value)] = true + } + return result +} + +func nullStringPtr(value sql.NullString) *string { + if !value.Valid { + return nil + } + return &value.String +} + +func nullIntPtr(value sql.NullInt64) *int { + if !value.Valid { + return nil + } + converted := int(value.Int64) + return &converted +} diff --git a/agents/drivers/vastbase/build.gradle b/agents/drivers/vastbase/build.gradle deleted file mode 100644 index 11ece4ac1..000000000 --- a/agents/drivers/vastbase/build.gradle +++ /dev/null @@ -1,10 +0,0 @@ -dependencies { - implementation fileTree(dir: 'libs', include: ['*.jar']) - implementation 'cn.com.vastdata:vastbase-jdbc:2.11v' -} - -tasks.named('shadowJar') { - manifest { - attributes('Agent-Label': 'Vastbase', 'Main-Class': 'com.dbx.agent.vastbase.VastbaseAgent') - } -} diff --git a/agents/drivers/vastbase/libs/.gitkeep b/agents/drivers/vastbase/libs/.gitkeep deleted file mode 100644 index e69de29bb..000000000 diff --git a/agents/drivers/vastbase/src/main/java/com/dbx/agent/vastbase/VastbaseAgent.java b/agents/drivers/vastbase/src/main/java/com/dbx/agent/vastbase/VastbaseAgent.java deleted file mode 100644 index 6ebe37fc8..000000000 --- a/agents/drivers/vastbase/src/main/java/com/dbx/agent/vastbase/VastbaseAgent.java +++ /dev/null @@ -1,61 +0,0 @@ -package com.dbx.agent.vastbase; - -import com.dbx.agent.MultiSessionJsonRpcServer; -import com.dbx.agent.ObjectSource; -import com.dbx.agent.PostgresLikeAgent; -import com.dbx.agent.PostgresLikeAgentProfile; - -public final class VastbaseAgent extends PostgresLikeAgent { - public static final PostgresLikeAgentProfile VASTBASE_PROFILE = new PostgresLikeAgentProfile( - "cn.com.vastbase.Driver", - "jdbc:vastbase://{host}:{port}/{database}" - ).withCatalogAttributeArraysMappedInJava(); - - public VastbaseAgent() { - super(VASTBASE_PROFILE); - } - - public static void main(String[] args) { - new MultiSessionJsonRpcServer(VastbaseAgent::new).run(); - } - - @Override - public ObjectSource getObjectSource(String schema, String name, String objectType) { - ObjectSource result = super.getObjectSource(schema, name, objectType); - String source = unwrapRecordText(result.getSource()); - return new ObjectSource(result.getName(), result.getObject_type(), result.getSchema(), source); - } - - /** - * Vastbase 的 pg_get_functiondef 在 SQL Server 兼容模式下返回的是 - * PostgreSQL 行记录文本格式 {@code (1,"source text")} 而非纯文本。 - * 此方法剥离行记录包装,提取实际的源码内容。 - */ - static String unwrapRecordText(String source) { - if (source == null || source.isEmpty()) { - return source; - } - String trimmed = source.trim(); - if (!trimmed.startsWith("(") || !trimmed.endsWith(")")) { - return source; - } - // 格式: (field1,field2,...),第一个字段通常是整数,第二个是源码文本 - String inner = trimmed.substring(1, trimmed.length() - 1); - int commaIdx = inner.indexOf(','); - if (commaIdx <= 0) { - // 单字段行记录: (source_text) - return unquoteRecordField(inner); - } - // 多字段行记录: 跳过第一个字段,提取第二个字段及之后的内容 - String rest = inner.substring(commaIdx + 1); - return unquoteRecordField(rest); - } - - private static String unquoteRecordField(String value) { - String trimmed = value.trim(); - if (trimmed.startsWith("\"") && trimmed.endsWith("\"")) { - return trimmed.substring(1, trimmed.length() - 1).replace("\"\"", "\""); - } - return trimmed; - } -} diff --git a/agents/drivers/vastbase/src/test/java/com/dbx/agent/vastbase/VastbaseAgentTest.java b/agents/drivers/vastbase/src/test/java/com/dbx/agent/vastbase/VastbaseAgentTest.java deleted file mode 100644 index 049766fff..000000000 --- a/agents/drivers/vastbase/src/test/java/com/dbx/agent/vastbase/VastbaseAgentTest.java +++ /dev/null @@ -1,521 +0,0 @@ -package com.dbx.agent.vastbase; - -import com.dbx.agent.AbstractJdbcAgent; -import com.dbx.agent.ColumnInfo; -import com.dbx.agent.DatabaseAgent; -import com.dbx.agent.ForeignKeyInfo; -import com.dbx.agent.IndexInfo; -import com.dbx.agent.test.JdbcFakeExecutionBehaviorTest; -import org.junit.jupiter.api.Assertions; -import org.junit.jupiter.api.Test; - -import java.lang.reflect.Field; -import java.lang.reflect.InvocationHandler; -import java.lang.reflect.Method; -import java.lang.reflect.Proxy; -import java.sql.Connection; -import java.sql.PreparedStatement; -import java.sql.ResultSet; -import java.util.ArrayList; -import java.util.Arrays; -import java.util.List; - -class VastbaseAgentTest extends JdbcFakeExecutionBehaviorTest { - @Override - protected DatabaseAgent createAgent() { - return new VastbaseAgent(); - } - - @Override - protected String resultSetSql() { - return "CALL sample_proc()"; - } - - @Test - void declaresVastbasePostgresLikeProfile() { - VastbaseAgent agent = new VastbaseAgent(); - - Assertions.assertEquals("cn.com.vastbase.Driver", agent.getProfile().getDriverClass()); - Assertions.assertEquals("jdbc:vastbase://{host}:{port}/{database}", agent.getProfile().getUrlTemplate()); - Assertions.assertTrue(agent.getProfile().mapsCatalogAttributeArraysInJava()); - } - - @Test - void mapsTablesWithoutPrimaryKeysFromCatalogArrays() throws Exception { - VastbaseMetadataFake metadata = new VastbaseMetadataFake(null); - VastbaseAgent agent = agentWithConnection(metadata.connection()); - - List columns = agent.getColumns("avatar_asset", "app_config_info"); - - Assertions.assertFalse(column(columns, "id").getIs_primary_key()); - Assertions.assertFalse(column(columns, "tenant_id").getIs_primary_key()); - metadata.assertUsesPostgres92CompatibleSql(); - } - - @Test - void mapsSinglePrimaryKeysFromCatalogArrays() throws Exception { - VastbaseMetadataFake metadata = new VastbaseMetadataFake(new Short[]{1}); - VastbaseAgent agent = agentWithConnection(metadata.connection()); - - List columns = agent.getColumns("avatar_asset", "app_config_info"); - - Assertions.assertTrue(column(columns, "id").getIs_primary_key()); - Assertions.assertFalse(column(columns, "tenant_id").getIs_primary_key()); - metadata.assertUsesPostgres92CompatibleSql(); - } - - @Test - void rejectsEntirePrimaryKeyWhenAnyAttributeNumberIsUnknown() throws Exception { - VastbaseMetadataFake metadata = new VastbaseMetadataFake(new Short[]{1, 99}); - VastbaseAgent agent = agentWithConnection(metadata.connection()); - - List columns = agent.getColumns("avatar_asset", "app_config_info"); - - Assertions.assertFalse(column(columns, "id").getIs_primary_key()); - Assertions.assertFalse(column(columns, "tenant_id").getIs_primary_key()); - Assertions.assertFalse(column(columns, "payload").getIs_primary_key()); - metadata.assertUsesPostgres92CompatibleSql(); - } - - @Test - void mapsCompositeKeysIndexesAndForeignKeysInCatalogOrder() throws Exception { - VastbaseMetadataFake metadata = new VastbaseMetadataFake(new Short[]{2, 1}); - VastbaseAgent agent = agentWithConnection(metadata.connection()); - - List columns = agent.getColumns("avatar_asset", "app_config_info"); - List indexes = agent.listIndexes("avatar_asset", "app_config_info"); - List foreignKeys = agent.listForeignKeys("avatar_asset", "app_config_info"); - metadata.assertSqlArraysFreed(); - metadata.clearRecordedOperations(); - String ddl = agent.getTableDdl("avatar_asset", "app_config_info"); - - Assertions.assertTrue(column(columns, "id").getIs_primary_key()); - Assertions.assertTrue(column(columns, "tenant_id").getIs_primary_key()); - Assertions.assertFalse(column(columns, "payload").getIs_primary_key()); - Assertions.assertEquals(Arrays.asList( - new IndexInfo("idx_payload", Arrays.asList("payload"), false, false, null, "btree", null, null), - new IndexInfo("idx_tenant_id", Arrays.asList("tenant_id", "id"), true, false, null, "btree", null, null), - new IndexInfo("idx_primitive", Arrays.asList("id", "payload"), false, false, null, "btree", null, null), - new IndexInfo("idx_object_array_null", Arrays.asList("tenant_id", "payload"), false, false, null, "btree", null, null), - new IndexInfo("idx_object_array_unsupported", Arrays.asList("id", "tenant_id"), false, false, null, "btree", null, null), - new IndexInfo("app_config_info_pkey", Arrays.asList("tenant_id", "id"), true, true, null, "btree", null, null) - ), indexes); - Assertions.assertEquals(Arrays.asList( - new ForeignKeyInfo("fk_account_region", "tenant_id", "accounts", "account_id"), - new ForeignKeyInfo("fk_account_region", "id", "accounts", "region_id"), - new ForeignKeyInfo("fk_owner", "id", "users", "user_id") - ), foreignKeys); - Assertions.assertTrue(ddl.contains("PRIMARY KEY (\"tenant_id\", \"id\")"), ddl); - Assertions.assertTrue( - ddl.contains("CONSTRAINT \"fk_account_region\" FOREIGN KEY (\"tenant_id\", \"id\") REFERENCES \"accounts\"(\"account_id\", \"region_id\")"), - ddl - ); - Assertions.assertEquals(1, occurrences(ddl, "CONSTRAINT \"fk_account_region\""), ddl); - Assertions.assertFalse(ddl.contains("fk_length_mismatch"), ddl); - Assertions.assertFalse(ddl.contains("fk_unknown_ref"), ddl); - Assertions.assertTrue( - ddl.contains("CREATE UNIQUE INDEX \"idx_tenant_id\" ON \"avatar_asset\".\"app_config_info\" USING btree (\"tenant_id\", \"id\")"), - ddl - ); - Assertions.assertEquals(8, metadata.queryCount(), metadata.sql()); - Assertions.assertEquals(1, metadata.tableAttributeQueryCount(), metadata.sql()); - Assertions.assertEquals(1, metadata.referencedAttributeQueryCount(), metadata.sql()); - Assertions.assertTrue(metadata.sql().contains("WHERE a.attrelid IN (?, ?)"), metadata.sql()); - metadata.assertSqlArraysFreed(); - metadata.assertUsesPostgres92CompatibleSql(); - Assertions.assertTrue(metadata.sql().contains("co.conkey AS column_numbers"), metadata.sql()); - Assertions.assertTrue(metadata.sql().contains("co.confkey AS ref_column_numbers"), metadata.sql()); - Assertions.assertTrue(metadata.sql().contains("ix.indkey AS column_numbers"), metadata.sql()); - } - - private static int occurrences(String value, String fragment) { - int result = 0; - int offset = 0; - while (true) { - int found = value.indexOf(fragment, offset); - if (found < 0) { - return result; - } - result += 1; - offset = found + fragment.length(); - } - } - - private static VastbaseAgent agentWithConnection(Connection connection) throws Exception { - VastbaseAgent agent = new VastbaseAgent(); - Field connectionField = AbstractJdbcAgent.class.getDeclaredField("connection"); - connectionField.setAccessible(true); - connectionField.set(agent, connection); - return agent; - } - - private static ColumnInfo column(List columns, String name) { - for (ColumnInfo column : columns) { - if (name.equals(column.getName())) { - return column; - } - } - throw new AssertionError("Missing column: " + name); - } - - private static final class VastbaseMetadataFake { - private final Object primaryKeyNumbers; - private final List statements = new ArrayList<>(); - private final List sqlArrays = new ArrayList<>(); - - private VastbaseMetadataFake(Object primaryKeyNumbers) { - this.primaryKeyNumbers = primaryKeyNumbers; - } - - private Connection connection() { - return proxy(Connection.class, new MethodHandler() { - @Override - public Object handle(Method method, Object[] args) { - String methodName = method.getName(); - if ("prepareStatement".equals(methodName)) { - String sql = (String) args[0]; - statements.add(sql); - return preparedStatement(sql); - } - if ("isClosed".equals(methodName)) { - return false; - } - if ("close".equals(methodName)) { - return null; - } - return defaultValue(method.getReturnType()); - } - }); - } - - private PreparedStatement preparedStatement(String sql) { - return proxy(PreparedStatement.class, new MethodHandler() { - @Override - public Object handle(Method method, Object[] args) { - String methodName = method.getName(); - if ("executeQuery".equals(methodName)) { - return resultFor(sql); - } - if ("setLong".equals(methodName)) { - return null; - } - if ("setString".equals(methodName) || "close".equals(methodName)) { - return null; - } - return defaultValue(method.getReturnType()); - } - }); - } - - private ResultSet resultFor(String sql) { - if (sql.contains("co.contype = 'p'")) { - if (primaryKeyNumbers == null) { - return resultSet(new String[]{"column_numbers"}, new Object[0][]); - } - return resultSet( - new String[]{"column_numbers"}, - new Object[][]{{primaryKeyNumbers}} - ); - } - if (sql.contains("ix.indkey AS column_numbers")) { - return resultSet( - new String[]{"index_name", "index_type", "is_unique", "is_primary", "column_numbers"}, - new Object[][]{ - {"idx_payload", "btree", false, false, sqlArray(new int[]{3})}, - {"idx_tenant_id", "btree", true, false, "2 1"}, - {"idx_primitive", "btree", false, false, new short[]{1, 3}}, - {"idx_object_array_null", "btree", false, false, objectSqlArray(new int[]{2, 3}, false)}, - {"idx_object_array_unsupported", "btree", false, false, objectSqlArray(new int[]{1, 2}, true)}, - {"app_config_info_pkey", "btree", true, true, new Short[]{2, 1}}, - {"idx_expression", "btree", false, false, new short[]{0, 3}}, - {"idx_unknown", "btree", false, false, "{2,99}"} - } - ); - } - if (sql.contains("co.confkey AS ref_column_numbers")) { - return resultSet( - new String[]{ - "constraint_name", - "column_numbers", - "ref_column_numbers", - "ref_table", - "ref_table_oid" - }, - new Object[][]{ - {"fk_account_region", "{2,1}", "{1,3}", "accounts", 42L}, - {"fk_length_mismatch", "{1,2}", "{1}", "accounts", 42L}, - {"fk_owner", new Short[]{1}, new Short[]{1}, "users", 43L}, - {"fk_unknown_ref", "{1}", "{99}", "users", 43L} - } - ); - } - if (sql.contains("WHERE a.attrelid IN (")) { - return resultSet( - new String[]{"relation_oid", "attribute_number", "column_name"}, - new Object[][]{ - {42L, 1, "account_id"}, - {42L, 3, "region_id"}, - {43L, 1, "user_id"} - } - ); - } - if (sql.contains("a.attnum AS attribute_number")) { - return attributeResult(new Object[][]{ - {1, "id"}, - {2, "tenant_id"}, - {3, "payload"} - }); - } - if (sql.contains("AS constraint_definition")) { - return resultSet( - new String[]{"constraint_name", "constraint_definition"}, - new Object[0][] - ); - } - if (sql.contains("AS table_comment")) { - return resultSet(new String[]{"table_comment"}, new Object[][]{{null}}); - } - if (sql.contains("AS data_type")) { - return resultSet( - new String[]{ - "column_name", - "data_type", - "is_nullable", - "column_default", - "column_comment", - "numeric_precision", - "numeric_scale", - "character_maximum_length" - }, - new Object[][]{ - {"id", "bigint", false, null, null, null, null, null}, - {"tenant_id", "bigint", false, null, null, null, null, null}, - {"payload", "text", true, null, null, null, null, null} - } - ); - } - throw new AssertionError("Unexpected SQL: " + sql); - } - - private ResultSet attributeResult(Object[][] rows) { - return resultSet(new String[]{"attribute_number", "column_name"}, rows); - } - - private String sql() { - return String.join("\n", statements); - } - - private java.sql.Array sqlArray(Object values) { - TrackingSqlArray sqlArray = new TrackingSqlArray(values); - sqlArrays.add(sqlArray); - return sqlArray.value(); - } - - private Object objectSqlArray(Object values, boolean unsupported) { - TrackingSqlArray sqlArray = new TrackingSqlArray(values); - sqlArrays.add(sqlArray); - return new ObjectSqlArray(sqlArray.value(), unsupported); - } - - private void clearRecordedOperations() { - statements.clear(); - sqlArrays.clear(); - } - - private int queryCount() { - return statements.size(); - } - - private int tableAttributeQueryCount() { - int result = 0; - for (String statement : statements) { - if (statement.contains("a.attnum AS attribute_number") - && statement.contains("JOIN pg_catalog.pg_class c")) { - result += 1; - } - } - return result; - } - - private int referencedAttributeQueryCount() { - int result = 0; - for (String statement : statements) { - if (statement.contains("a.attrelid AS relation_oid")) { - result += 1; - } - } - return result; - } - - private void assertSqlArraysFreed() { - Assertions.assertFalse(sqlArrays.isEmpty()); - for (TrackingSqlArray sqlArray : sqlArrays) { - Assertions.assertTrue(sqlArray.freed); - } - } - - private void assertUsesPostgres92CompatibleSql() { - Assertions.assertFalse(sql().contains("LATERAL"), sql()); - Assertions.assertFalse(sql().contains("WITH ORDINALITY"), sql()); - } - } - - private static ResultSet resultSet(String[] columns, Object[][] rows) { - int[] rowIndex = {-1}; - boolean[] wasNull = {false}; - return proxy(ResultSet.class, new MethodHandler() { - @Override - public Object handle(Method method, Object[] args) { - String methodName = method.getName(); - if ("next".equals(methodName)) { - rowIndex[0] += 1; - return rowIndex[0] < rows.length; - } - if ("getString".equals(methodName)) { - Object value = columnValue(columns, rows[rowIndex[0]], args[0]); - wasNull[0] = value == null; - return value == null ? null : String.valueOf(value); - } - if ("getBoolean".equals(methodName)) { - Object value = columnValue(columns, rows[rowIndex[0]], args[0]); - wasNull[0] = value == null; - return value instanceof Boolean && (Boolean) value; - } - if ("getInt".equals(methodName)) { - Object value = columnValue(columns, rows[rowIndex[0]], args[0]); - wasNull[0] = value == null; - return value == null ? 0 : ((Number) value).intValue(); - } - if ("getLong".equals(methodName)) { - Object value = columnValue(columns, rows[rowIndex[0]], args[0]); - wasNull[0] = value == null; - return value == null ? 0L : ((Number) value).longValue(); - } - if ("getObject".equals(methodName)) { - Object value = columnValue(columns, rows[rowIndex[0]], args[0]); - if (value instanceof ObjectSqlArray) { - value = ((ObjectSqlArray) value).value; - } - wasNull[0] = value == null; - return value; - } - if ("getArray".equals(methodName)) { - Object value = columnValue(columns, rows[rowIndex[0]], args[0]); - if (value instanceof ObjectSqlArray) { - ObjectSqlArray objectSqlArray = (ObjectSqlArray) value; - if (objectSqlArray.unsupported) { - throw new UnsupportedOperationException(); - } - return null; - } - wasNull[0] = value == null; - return value instanceof java.sql.Array ? value : null; - } - if ("wasNull".equals(methodName)) { - return wasNull[0]; - } - if ("close".equals(methodName)) { - return null; - } - return defaultValue(method.getReturnType()); - } - }); - } - - private static Object columnValue(String[] columns, Object[] row, Object key) { - if (key instanceof Number) { - return row[((Number) key).intValue() - 1]; - } - for (int columnIndex = 0; columnIndex < columns.length; columnIndex++) { - if (columns[columnIndex].equalsIgnoreCase(String.valueOf(key))) { - return row[columnIndex]; - } - } - return null; - } - - private static T proxy(Class type, MethodHandler handler) { - InvocationHandler invocationHandler = new InvocationHandler() { - @Override - public Object invoke(Object proxy, Method method, Object[] args) { - return handler.handle(method, args); - } - }; - return type.cast(Proxy.newProxyInstance( - type.getClassLoader(), - new Class[]{type}, - invocationHandler - )); - } - - private static Object defaultValue(Class type) { - if (Boolean.TYPE.equals(type)) { - return false; - } - if (Byte.TYPE.equals(type)) { - return (byte) 0; - } - if (Short.TYPE.equals(type)) { - return (short) 0; - } - if (Integer.TYPE.equals(type)) { - return 0; - } - if (Long.TYPE.equals(type)) { - return 0L; - } - if (Float.TYPE.equals(type)) { - return 0.0f; - } - if (Double.TYPE.equals(type)) { - return 0.0; - } - return null; - } - - private static final class TrackingSqlArray { - private final Object values; - private boolean freed; - - private TrackingSqlArray(Object values) { - this.values = values; - } - - private java.sql.Array value() { - return proxy(java.sql.Array.class, new MethodHandler() { - @Override - public Object handle(Method method, Object[] args) { - String methodName = method.getName(); - if ("getArray".equals(methodName) && (args == null || args.length == 0)) { - return values; - } - if ("free".equals(methodName)) { - freed = true; - return null; - } - if ("getBaseType".equals(methodName)) { - return java.sql.Types.SMALLINT; - } - if ("getBaseTypeName".equals(methodName)) { - return "int2"; - } - return defaultValue(method.getReturnType()); - } - }); - } - } - - private static final class ObjectSqlArray { - private final java.sql.Array value; - private final boolean unsupported; - - private ObjectSqlArray(java.sql.Array value, boolean unsupported) { - this.value = value; - this.unsupported = unsupported; - } - } - - private interface MethodHandler { - Object handle(Method method, Object[] args); - } -} diff --git a/agents/metadata-constraint-coverage.tsv b/agents/metadata-constraint-coverage.tsv index 8e806bd41..b2120630a 100644 --- a/agents/metadata-constraint-coverage.tsv +++ b/agents/metadata-constraint-coverage.tsv @@ -12,6 +12,7 @@ h2 native-pushdown java-sql Uses INFORMATION_SCHEMA with type, filter, stable or hive intentional-fallback java-sql Uses Hive metadata paths where stable server-side fuzzy paging is not portable, common constraints filter locally. informix native-pushdown java-sql Uses systables/sysprocedures with type, filter, stable order, and SKIP/FIRST. kingbase-go shared-fallback native-go Uses sys_catalog/information_schema for compatibility-aware metadata, then applies stable filtering and paging in the native agent. +vastbase-go shared-fallback native-go Uses pg_catalog/information_schema for Vastbase metadata, then applies stable filtering and paging in the native agent. kylin intentional-fallback java-jdbc-metadata Uses JDBC metadata from Kylin driver; no portable server-side paging API, common constraints filter locally. neo4j intentional-fallback java-graph-metadata Uses JDBC metadata with Cypher fallback; label discovery cannot guarantee relational-style filtered paging, common constraints filter locally. oceanbase-oracle native-pushdown java-sql Uses Oracle-compatible metadata SQL with type, filter, stable order, and ROWNUM paging. diff --git a/agents/scripts/driver_release_packages_test.py b/agents/scripts/driver_release_packages_test.py index 148eee9ef..b96fcef03 100644 --- a/agents/scripts/driver_release_packages_test.py +++ b/agents/scripts/driver_release_packages_test.py @@ -18,6 +18,8 @@ class DriverReleasePackagesTest(unittest.TestCase): release_dir = Path(temp_dir) native_source = release_dir / "dbx-agent-kingbase-windows-x64.exe" native_source.write_bytes(b"MZtest-agent") + vastbase_source = release_dir / "dbx-agent-vastbase-linux-x64" + vastbase_source.write_bytes(b"\x7fELFtest-vastbase-agent") duckdb_source = release_dir / "dbx-agent-duckdb-macos-aarch64" duckdb_source.write_bytes(b"\xcf\xfa\xed\xfetest-duckdb-agent") rabbitmq_source = release_dir / "dbx-agent-rabbitmq-linux-x64" @@ -29,6 +31,7 @@ class DriverReleasePackagesTest(unittest.TestCase): "oracle": "0.1.10", "xugu": "0.1.20", "kingbase": "0.1.34", + "vastbase": "0.1.37", "duckdb": "0.1.0", "rabbitmq": "0.1.0", } @@ -36,9 +39,10 @@ class DriverReleasePackagesTest(unittest.TestCase): renamed = version_agent_artifacts(release_dir, versions) versioned_java = release_dir / "dbx-agent-h2-0.2.5.jar" versioned_native = release_dir / "dbx-agent-kingbase-0.1.34-windows-x64.exe" + versioned_vastbase = release_dir / "dbx-agent-vastbase-0.1.37-linux-x64" versioned_duckdb = release_dir / "dbx-agent-duckdb-0.1.0-macos-aarch64" versioned_rabbitmq = release_dir / "dbx-agent-rabbitmq-0.1.0-linux-x64" - self.assertEqual(renamed, [versioned_java, versioned_native, versioned_duckdb, versioned_rabbitmq]) + self.assertEqual(renamed, [versioned_java, versioned_native, versioned_vastbase, versioned_duckdb, versioned_rabbitmq]) registry = { "jres": {"21": {"version": "21", "platforms": {}}}, @@ -63,6 +67,19 @@ class DriverReleasePackagesTest(unittest.TestCase): } }, }, + "vastbase": { + "version": "0.1.37", + "label": "Vastbase", + "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_vastbase.name}", + "size": versioned_vastbase.stat().st_size, + } + }, + }, "duckdb": { "version": "0.1.0", "label": "DuckDB", @@ -100,6 +117,7 @@ class DriverReleasePackagesTest(unittest.TestCase): [ release_dir / "dbx-agent-h2-0.2.5.tar.zst", release_dir / "dbx-agent-kingbase-0.1.34-windows-x64.tar.zst", + release_dir / "dbx-agent-vastbase-0.1.37-linux-x64.tar.zst", release_dir / "dbx-agent-duckdb-0.1.0-macos-aarch64.tar.zst", release_dir / "dbx-agent-rabbitmq-0.1.0-linux-x64.tar.zst", ], @@ -107,8 +125,9 @@ class DriverReleasePackagesTest(unittest.TestCase): package_cases = [ (outputs[0], "h2", versioned_java, "jar", None), (outputs[1], "kingbase", versioned_native, "native", "windows-x64"), - (outputs[2], "duckdb", versioned_duckdb, "native", "macos-aarch64"), - (outputs[3], "rabbitmq", versioned_rabbitmq, "native", "linux-x64"), + (outputs[2], "vastbase", versioned_vastbase, "native", "linux-x64"), + (outputs[3], "duckdb", versioned_duckdb, "native", "macos-aarch64"), + (outputs[4], "rabbitmq", versioned_rabbitmq, "native", "linux-x64"), ] for output, driver_name, source, artifact_type, platform in package_cases: tar_bytes = subprocess.run( @@ -134,8 +153,9 @@ class DriverReleasePackagesTest(unittest.TestCase): release_artifacts = [ (final_registry["drivers"]["h2"]["jar"], outputs[0]), (final_registry["drivers"]["kingbase"]["native"]["windows-x64"], outputs[1]), - (final_registry["drivers"]["duckdb"]["native"]["macos-aarch64"], outputs[2]), - (final_registry["drivers"]["rabbitmq"]["native"]["linux-x64"], outputs[3]), + (final_registry["drivers"]["vastbase"]["native"]["linux-x64"], outputs[2]), + (final_registry["drivers"]["duckdb"]["native"]["macos-aarch64"], outputs[3]), + (final_registry["drivers"]["rabbitmq"]["native"]["linux-x64"], outputs[4]), ] for artifact, output in release_artifacts: self.assertEqual(artifact["url"], f"https://example.com/{output.name}") @@ -144,7 +164,7 @@ class DriverReleasePackagesTest(unittest.TestCase): self.assertEqual(len(artifact["sha256"]), 64) removed = remove_raw_driver_artifacts(release_dir) - self.assertEqual(removed, [versioned_duckdb, versioned_java, versioned_native, versioned_rabbitmq]) + self.assertEqual(removed, [versioned_duckdb, versioned_java, versioned_native, versioned_rabbitmq, versioned_vastbase]) self.assertTrue(all(output.is_file() for output in outputs)) def test_full_offline_bundle_includes_supported_windows_artifacts(self) -> None: diff --git a/agents/scripts/validate_agents.py b/agents/scripts/validate_agents.py index 02cc61d52..fa7526f97 100644 --- a/agents/scripts/validate_agents.py +++ b/agents/scripts/validate_agents.py @@ -16,6 +16,7 @@ NATIVE_ONLY_AGENT_MODULES = { "duckdb": "drivers/duckdb", "oracle": "drivers/oracle-go", "kingbase": "drivers/kingbase-go", + "vastbase": "drivers/vastbase-go", "xugu": "drivers/xugu", "rabbitmq": "drivers/rabbitmq", } diff --git a/agents/scripts/validate_agents_test.py b/agents/scripts/validate_agents_test.py index e1be4bfdb..e946d8db3 100644 --- a/agents/scripts/validate_agents_test.py +++ b/agents/scripts/validate_agents_test.py @@ -136,10 +136,10 @@ class ValidateAgentsTest(unittest.TestCase): "include(*(infrastructureModules + driverModules))\n", encoding="utf-8", ) - for driver in ("oracle-go", "kingbase-go", "xugu", "duckdb", "rabbitmq"): + for driver in ("oracle-go", "kingbase-go", "vastbase-go", "xugu", "duckdb", "rabbitmq"): (root / "drivers" / driver).mkdir(parents=True) (root / "versions.json").write_text( - json.dumps({"h2": "0.1.0", "oracle": "0.1.0", "kingbase": "0.1.0", "xugu": "0.1.0", "rabbitmq": "0.1.0"}), + json.dumps({"h2": "0.1.0", "oracle": "0.1.0", "kingbase": "0.1.0", "vastbase": "0.1.0", "xugu": "0.1.0", "rabbitmq": "0.1.0"}), encoding="utf-8", ) diff --git a/agents/scripts/version_agent_artifacts.py b/agents/scripts/version_agent_artifacts.py index ebbce01cd..e1e8b1a60 100644 --- a/agents/scripts/version_agent_artifacts.py +++ b/agents/scripts/version_agent_artifacts.py @@ -4,7 +4,7 @@ import json from pathlib import Path -NATIVE_DRIVERS = ("oracle", "xugu", "kingbase", "duckdb", "rabbitmq") +NATIVE_DRIVERS = ("oracle", "xugu", "kingbase", "vastbase", "duckdb", "rabbitmq") PLATFORMS = ( "macos-aarch64", "macos-x64", diff --git a/agents/settings.gradle b/agents/settings.gradle index 12e530ea1..78a7f3b64 100644 --- a/agents/settings.gradle +++ b/agents/settings.gradle @@ -2,7 +2,7 @@ rootProject.name = 'dbx-agents' def infrastructureModules = ['common', 'test-support'] def driverModules = [ - 'access', 'dameng', 'vastbase', 'goldendb', 'databend', 'databricks', 'saphana', + 'access', 'dameng', 'goldendb', 'databend', 'databricks', 'saphana', 'teradata', 'vertica', 'firebird', 'exasol', 'oceanbase-oracle', 'gbase8a', 'gbase8s', 'bigquery', 'kylin', 'sundb', 'h2', 'h2-legacy', 'snowflake', 'trino', 'hive', 'spark', 'db2', 'informix', 'neo4j', 'cassandra', 'mongodb', 'highgo', 'uxdb', 'tdengine', 'yashandb', 'oscar', diff --git a/packages/app-tests/agentVersionBump.test.ts b/packages/app-tests/agentVersionBump.test.ts index 22a497b79..59dd344d5 100644 --- a/packages/app-tests/agentVersionBump.test.ts +++ b/packages/app-tests/agentVersionBump.test.ts @@ -150,6 +150,23 @@ test("Kingbase native Go source changes bump the Kingbase module version", () => ]); }); +test("Vastbase native Go source changes bump the Vastbase module version", () => { + const fixture = moduleFixture(["agents/drivers/vastbase-go"]); + + const result = evaluateAgentVersionBump({ + versions: { + vastbase: "0.1.37", + }, + changedFiles: ["agents/drivers/vastbase-go/main.go"], + ...fixture, + }); + + assert.equal(result.changed, true); + assert.deepEqual(result.versions, { + vastbase: "0.1.38", + }); +}); + test("manual agent versions are preserved while other changed modules auto bump", () => { const fixture = moduleFixture([ "agents/drivers/access",