feat(vastbase): replace JDBC agent with native Go driver

This commit is contained in:
t8y2 2026-08-03 23:08:52 +08:00
parent 05841277d1
commit dae150ea74
No known key found for this signature in database
42 changed files with 6528 additions and 618 deletions

View File

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

View File

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

View File

@ -58,6 +58,7 @@ const DRIVER_DATABASE_ALIASES = {
rabbitmq: "mq",
rocketmq: "mq",
"sqlserver-legacy": "sqlserver",
"vastbase-go": "vastbase",
};
const DIALECT_DATABASE_ALIASES = {

View File

@ -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"],
);
});

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@ -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 运行时测试

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

File diff suppressed because it is too large Load Diff

View File

@ -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&currentSchema=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
}

View File

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

View File

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

View File

@ -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[:])
}

View File

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

View File

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

View File

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

File diff suppressed because it is too large Load Diff

View File

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

View File

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

View File

@ -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<ColumnInfo> 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<ColumnInfo> 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<ColumnInfo> 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<ColumnInfo> columns = agent.getColumns("avatar_asset", "app_config_info");
List<IndexInfo> indexes = agent.listIndexes("avatar_asset", "app_config_info");
List<ForeignKeyInfo> 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<ColumnInfo> 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<String> statements = new ArrayList<>();
private final List<TrackingSqlArray> 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> T proxy(Class<T> 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);
}
}

View File

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

1 driver strategy scope reason
12 hive intentional-fallback java-sql Uses Hive metadata paths where stable server-side fuzzy paging is not portable, common constraints filter locally.
13 informix native-pushdown java-sql Uses systables/sysprocedures with type, filter, stable order, and SKIP/FIRST.
14 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.
15 vastbase-go shared-fallback native-go Uses pg_catalog/information_schema for Vastbase metadata, then applies stable filtering and paging in the native agent.
16 kylin intentional-fallback java-jdbc-metadata Uses JDBC metadata from Kylin driver; no portable server-side paging API, common constraints filter locally.
17 neo4j intentional-fallback java-graph-metadata Uses JDBC metadata with Cypher fallback; label discovery cannot guarantee relational-style filtered paging, common constraints filter locally.
18 oceanbase-oracle native-pushdown java-sql Uses Oracle-compatible metadata SQL with type, filter, stable order, and ROWNUM paging.

View File

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

View File

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

View File

@ -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",
)

View File

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

View File

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

View File

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