feat(vastbase): replace JDBC agent with native Go driver
This commit is contained in:
parent
05841277d1
commit
dae150ea74
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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" },
|
||||
|
|
|
|||
|
|
@ -58,6 +58,7 @@ const DRIVER_DATABASE_ALIASES = {
|
|||
rabbitmq: "mq",
|
||||
rocketmq: "mq",
|
||||
"sqlserver-legacy": "sqlserver",
|
||||
"vastbase-go": "vastbase",
|
||||
};
|
||||
|
||||
const DIALECT_DATABASE_ALIASES = {
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
);
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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/,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 运行时测试
|
||||
|
||||
|
|
|
|||
|
|
@ -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) }
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
@ -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.
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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(' ')
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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=
|
||||
|
|
@ -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
|
|
@ -0,0 +1,281 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
pq "gitcode.com/opengauss/openGauss-connector-go-pq"
|
||||
)
|
||||
|
||||
func TestVastbaseHandshakeAdvertisesMultiSessionSQLAgent(t *testing.T) {
|
||||
runtime := &runtimeServer{sessions: map[string]*agentSession{}}
|
||||
result, shutdown, err := runtime.dispatch("handshake", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("handshake failed: %v", err)
|
||||
}
|
||||
if shutdown {
|
||||
t.Fatal("handshake must not request shutdown")
|
||||
}
|
||||
payload, err := json.Marshal(result)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal handshake: %v", err)
|
||||
}
|
||||
text := string(payload)
|
||||
for _, expected := range []string{`"protocolVersion":2`, `"multi_session"`, `"metadata"`, `"paged_query"`, `"structured_error_v1"`} {
|
||||
if !strings.Contains(text, expected) {
|
||||
t.Fatalf("handshake missing %s: %s", expected, text)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestQueryOptionsFromParams(t *testing.T) {
|
||||
params := map[string]json.RawMessage{
|
||||
"sql": json.RawMessage(`"SELECT 1"`),
|
||||
"database": json.RawMessage(`"dbx"`),
|
||||
"schema": json.RawMessage(`"public"`),
|
||||
"maxRows": json.RawMessage(`1000`),
|
||||
"fetchSize": json.RawMessage(`250`),
|
||||
"timeoutSecs": json.RawMessage(`15`),
|
||||
}
|
||||
expected := queryOptions{SQL: "SELECT 1", Database: "dbx", Schema: "public", MaxRows: 1000, FetchSize: 250, TimeoutSecs: 15}
|
||||
if actual := queryOptionsFromParams(params); actual != expected {
|
||||
t.Fatalf("queryOptionsFromParams() = %+v, want %+v", actual, expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVastbaseBuildDSNUsesNativeDefaultsForJDBCURL(t *testing.T) {
|
||||
dsn := buildDSN(connectParams{
|
||||
Host: "vastbase.example.com",
|
||||
Database: "postgres",
|
||||
Username: "vbadmin",
|
||||
Password: "secret",
|
||||
ConnectionString: "jdbc:vastbase://vastbase.example.com:5432/postgres",
|
||||
URLParams: "application_name=dbx",
|
||||
})
|
||||
for _, expected := range []string{
|
||||
"host='vastbase.example.com'",
|
||||
"port=5432",
|
||||
"user='vbadmin'",
|
||||
"password='secret'",
|
||||
"dbname='postgres'",
|
||||
"sslmode=prefer",
|
||||
"application_name='dbx'",
|
||||
} {
|
||||
if !strings.Contains(dsn, expected) {
|
||||
t.Fatalf("DSN missing %s: %s", expected, dsn)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestVastbaseBuildDSNPreservesNativeConnectionString(t *testing.T) {
|
||||
dsn := buildDSNWithSSLMode(connectParams{
|
||||
ConnectionString: "postgresql://vbadmin:secret@vastbase.example.com:5432/postgres?application_name=dbx&sslmode=disable",
|
||||
}, "verify-full")
|
||||
if !strings.Contains(dsn, "application_name=dbx") || !strings.Contains(dsn, "sslmode=verify-full") {
|
||||
t.Fatalf("unexpected rewritten native DSN: %s", dsn)
|
||||
}
|
||||
if strings.Contains(dsn, "sslmode=disable") {
|
||||
t.Fatalf("old sslmode was not replaced: %s", dsn)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVastbaseBuildDSNTranslatesJDBCParameters(t *testing.T) {
|
||||
dsn := buildDSN(connectParams{
|
||||
Host: "vastbase.example.com",
|
||||
Database: "postgres",
|
||||
Username: "vbadmin",
|
||||
Password: "secret",
|
||||
URLParams: "targetServerType=master&connectTimeout=7¤tSchema=app&applicationName=dbx&sslmode=enable&autosave=always&enable_ce=1&db_compatibility=PG",
|
||||
})
|
||||
for _, expected := range []string{
|
||||
"target_session_attrs='primary'",
|
||||
"connect_timeout='7'",
|
||||
"search_path='app'",
|
||||
"application_name='dbx'",
|
||||
"sslmode=require",
|
||||
} {
|
||||
if !strings.Contains(dsn, expected) {
|
||||
t.Fatalf("translated DSN missing %s: %s", expected, dsn)
|
||||
}
|
||||
}
|
||||
for _, rejected := range []string{"targetServerType", "currentSchema", "applicationName", "autosave", "enable_ce", "db_compatibility"} {
|
||||
if strings.Contains(dsn, rejected) {
|
||||
t.Fatalf("JDBC-only parameter leaked into DSN: %s", dsn)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestVastbaseDriverRegistration(t *testing.T) {
|
||||
if !containsString(sql.Drivers(), agentSQLDriverName) {
|
||||
t.Fatalf("%s driver is not registered: %v", agentSQLDriverName, sql.Drivers())
|
||||
}
|
||||
}
|
||||
|
||||
func TestVastbaseObjectSourceNormalization(t *testing.T) {
|
||||
tests := map[string]string{
|
||||
`(1,"CREATE FUNCTION f() RETURNS int AS ''SELECT 1'';")`: `CREATE FUNCTION f() RETURNS int AS ''SELECT 1'';`,
|
||||
`("CREATE VIEW v AS SELECT 1")`: `CREATE VIEW v AS SELECT 1`,
|
||||
`CREATE VIEW v AS SELECT 1`: `CREATE VIEW v AS SELECT 1`,
|
||||
}
|
||||
for input, expected := range tests {
|
||||
if actual := normalizeAgentObjectSource(input); actual != expected {
|
||||
t.Fatalf("normalizeAgentObjectSource(%q) = %q, want %q", input, actual, expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestVastbaseDataTypesIncludeVectorFamilies(t *testing.T) {
|
||||
types := agentDataTypes()
|
||||
for _, expected := range []string{"floatvector", "halfvector", "int8vector", "sparsevector"} {
|
||||
if !containsString(types, expected) {
|
||||
t.Fatalf("missing Vastbase data type %s: %v", expected, types)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestVastbaseMetadataErrorClassificationUsesOpenGaussCodes(t *testing.T) {
|
||||
undefinedColumn := &pq.Error{Code: pq.ErrorCode("42703"), Message: "column a.attidentity does not exist"}
|
||||
if !isUndefinedColumn(undefinedColumn, "attidentity") {
|
||||
t.Fatal("undefined Vastbase column was not recognized")
|
||||
}
|
||||
undefinedFunction := &pq.Error{Code: pq.ErrorCode("42883"), Message: "function pg_get_expr does not exist"}
|
||||
if !isUndefinedFunction(undefinedFunction, "pg_get_expr") {
|
||||
t.Fatal("undefined Vastbase function was not recognized")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVastbaseModeUsesPostgresCatalog(t *testing.T) {
|
||||
mode := detectAgentMode(nil, false)
|
||||
if mode.compatibilityMode != "postgres" || !mode.postgresCatalog || mode.mysqlCompat {
|
||||
t.Fatalf("unexpected default Vastbase mode: %+v", mode)
|
||||
}
|
||||
mysqlMode := detectAgentMode(nil, true)
|
||||
if mysqlMode.compatibilityMode != "mysql" || !mysqlMode.postgresCatalog || !mysqlMode.mysqlCompat {
|
||||
t.Fatalf("unexpected MySQL-compatible Vastbase mode: %+v", mysqlMode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaConnectionRecoversAfterCanceledDriverConnection(t *testing.T) {
|
||||
registerVastbaseSchemaRetryDriver.Do(func() {
|
||||
sql.Register("vastbase-schema-retry-test", &schemaRetryDriver{})
|
||||
})
|
||||
schemaRetryOpens.Store(0)
|
||||
db, err := sql.Open("vastbase-schema-retry-test", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
db.SetMaxOpenConns(1)
|
||||
db.SetMaxIdleConns(1)
|
||||
|
||||
server := newServer()
|
||||
server.db = db
|
||||
conn, err := server.schemaConn(context.Background(), "public")
|
||||
if err != nil {
|
||||
t.Fatalf("schema connection did not recover: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
if opens := schemaRetryOpens.Load(); opens != 2 {
|
||||
t.Fatalf("expected one replacement connection, opened %d", opens)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateConnectionRecoversAfterCanceledDriverConnection(t *testing.T) {
|
||||
registerVastbasePingRetryDriver.Do(func() {
|
||||
sql.Register("vastbase-ping-retry-test", &pingRetryDriver{})
|
||||
})
|
||||
pingRetryOpens.Store(0)
|
||||
db, err := sql.Open("vastbase-ping-retry-test", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
db.SetMaxOpenConns(1)
|
||||
db.SetMaxIdleConns(1)
|
||||
|
||||
server := newServer()
|
||||
server.db = db
|
||||
if err := server.validateConnection(); err != nil {
|
||||
t.Fatalf("connection validation did not recover: %v", err)
|
||||
}
|
||||
if opens := pingRetryOpens.Load(); opens != 2 {
|
||||
t.Fatalf("expected one replacement connection, opened %d", opens)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisconnectResetsInformationSchemaCapabilityCache(t *testing.T) {
|
||||
server := newServer()
|
||||
server.infoColumnTypeUnsupported = true
|
||||
server.infoUdtNameUnsupported = true
|
||||
|
||||
if err := server.disconnect(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if server.infoColumnTypeUnsupported || server.infoUdtNameUnsupported {
|
||||
t.Fatal("disconnect must reset cached information_schema capabilities")
|
||||
}
|
||||
}
|
||||
|
||||
var (
|
||||
registerVastbaseSchemaRetryDriver sync.Once
|
||||
registerVastbasePingRetryDriver sync.Once
|
||||
schemaRetryOpens atomic.Int32
|
||||
pingRetryOpens atomic.Int32
|
||||
)
|
||||
|
||||
type schemaRetryDriver struct{}
|
||||
|
||||
func (*schemaRetryDriver) Open(string) (driver.Conn, error) {
|
||||
return &schemaRetryConn{bad: schemaRetryOpens.Add(1) == 1}, nil
|
||||
}
|
||||
|
||||
type schemaRetryConn struct {
|
||||
bad bool
|
||||
}
|
||||
|
||||
func (*schemaRetryConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip }
|
||||
func (*schemaRetryConn) Close() error { return nil }
|
||||
func (*schemaRetryConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip }
|
||||
|
||||
func (conn *schemaRetryConn) ExecContext(context.Context, string, []driver.NamedValue) (driver.Result, error) {
|
||||
if conn.bad {
|
||||
return nil, driver.ErrBadConn
|
||||
}
|
||||
return driver.RowsAffected(0), nil
|
||||
}
|
||||
|
||||
type pingRetryDriver struct{}
|
||||
|
||||
func (*pingRetryDriver) Open(string) (driver.Conn, error) {
|
||||
return &pingRetryConn{bad: pingRetryOpens.Add(1) == 1}, nil
|
||||
}
|
||||
|
||||
type pingRetryConn struct {
|
||||
bad bool
|
||||
}
|
||||
|
||||
func (*pingRetryConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip }
|
||||
func (*pingRetryConn) Close() error { return nil }
|
||||
func (*pingRetryConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip }
|
||||
|
||||
func (conn *pingRetryConn) Ping(context.Context) error {
|
||||
if conn.bad {
|
||||
return driver.ErrBadConn
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func containsString(values []string, expected string) bool {
|
||||
for _, value := range values {
|
||||
if value == expected {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
|
@ -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()
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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[:])
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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
|
|
@ -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')
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
}
|
||||
}
|
||||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Reference in New Issue