feat(kingbase): add native Go agent
This commit is contained in:
parent
f5c614804a
commit
35c33a4085
|
|
@ -87,7 +87,9 @@ jobs:
|
|||
- uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: agent-jars
|
||||
path: "agents/drivers/*/build/libs/dbx-agent-*.jar"
|
||||
path: |
|
||||
agents/drivers/*/build/libs/dbx-agent-*.jar
|
||||
!agents/drivers/kingbase/build/libs/dbx-agent-kingbase.jar
|
||||
|
||||
build-oracle-native:
|
||||
needs: [bump-versions]
|
||||
|
|
@ -167,6 +169,45 @@ jobs:
|
|||
name: xugu-native
|
||||
path: "release-native/dbx-agent-xugu-*"
|
||||
|
||||
build-kingbase-native:
|
||||
needs: [bump-versions]
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version: "1.22.x"
|
||||
- name: Test Kingbase native agent
|
||||
working-directory: agents/drivers/kingbase-go
|
||||
run: go test ./...
|
||||
- name: Cross-compile Kingbase native agent
|
||||
shell: bash
|
||||
run: |
|
||||
mkdir -p release-native
|
||||
cd agents/drivers/kingbase-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-kingbase-${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: kingbase-native
|
||||
path: "release-native/dbx-agent-kingbase-*"
|
||||
|
||||
build-jre:
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
|
|
@ -239,7 +280,7 @@ jobs:
|
|||
path: "dbx-jre-*.tar.gz"
|
||||
|
||||
release:
|
||||
needs: [bump-versions, build-agents, build-oracle-native, build-xugu-native, build-jre]
|
||||
needs: [bump-versions, build-agents, build-oracle-native, build-xugu-native, build-kingbase-native, build-jre]
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
|
@ -250,9 +291,10 @@ jobs:
|
|||
- name: Flatten artifacts
|
||||
run: |
|
||||
mkdir -p release
|
||||
find artifacts/agent-jars -name '*.jar' -exec cp {} release/ \;
|
||||
find artifacts/agent-jars -name '*.jar' ! -name 'dbx-agent-kingbase.jar' -exec cp {} release/ \;
|
||||
find artifacts/oracle-native -type f -name 'dbx-agent-oracle-*' -exec cp {} release/ \;
|
||||
find artifacts/xugu-native -type f -name 'dbx-agent-xugu-*' -exec cp {} release/ \;
|
||||
find artifacts/kingbase-native -type f -name 'dbx-agent-kingbase-*' -exec cp {} release/ \;
|
||||
find artifacts -name 'dbx-jre-*.tar.gz' -exec cp {} release/ \;
|
||||
ls -lh release/
|
||||
|
||||
|
|
@ -312,6 +354,7 @@ jobs:
|
|||
native_only_label() {
|
||||
local name="$1"
|
||||
case "$name" in
|
||||
kingbase) echo "人大金仓 KingbaseES" ;;
|
||||
xugu) echo "虚谷 XuguDB" ;;
|
||||
*) echo "$name" ;;
|
||||
esac
|
||||
|
|
@ -370,7 +413,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; do
|
||||
for name in oracle xugu kingbase; do
|
||||
[ -f "release/dbx-agent-${name}.jar" ] && continue
|
||||
native_json=$(generate_native_platforms "$name")
|
||||
[ -z "$native_json" ] && continue
|
||||
|
|
@ -415,7 +458,7 @@ jobs:
|
|||
for f in release/dbx-agent-*.jar; do
|
||||
DRIVER_NAMES+=("$(basename "$f" .jar | sed 's/dbx-agent-//')")
|
||||
done
|
||||
for name in oracle xugu; do
|
||||
for name in oracle xugu kingbase; do
|
||||
[ -f "release/dbx-agent-${name}.jar" ] && continue
|
||||
compgen -G "release/dbx-agent-${name}-*" > /dev/null && DRIVER_NAMES+=("$name")
|
||||
done
|
||||
|
|
@ -423,6 +466,7 @@ jobs:
|
|||
native_only_label() {
|
||||
local name="$1"
|
||||
case "$name" in
|
||||
kingbase) echo "人大金仓 KingbaseES" ;;
|
||||
oracle) echo "Oracle" ;;
|
||||
xugu) echo "虚谷 XuguDB" ;;
|
||||
*) echo "$name" ;;
|
||||
|
|
@ -445,6 +489,8 @@ jobs:
|
|||
NOTES="${NOTES}### ${label} (${old_ver} → ${new_ver})"$'\n'
|
||||
if [ "$name" = "oracle" ]; then
|
||||
LOG_PATH="agents/drivers/oracle-go/"
|
||||
elif [ "$name" = "kingbase" ]; then
|
||||
LOG_PATH="agents/drivers/kingbase-go/"
|
||||
elif [ -d "agents/drivers/$name" ]; then
|
||||
LOG_PATH="agents/drivers/$name/"
|
||||
else
|
||||
|
|
|
|||
|
|
@ -8,11 +8,11 @@ Each agent runs as a standalone process and communicates with DBX via stdin/stdo
|
|||
|
||||
## Supported Databases
|
||||
|
||||
| Agent | Database | JDBC Driver |
|
||||
| Agent | Database | Driver |
|
||||
|-------|----------|-------------|
|
||||
| access | Microsoft Access | UCanAccess |
|
||||
| dameng | 达梦 DM8 | DM JDBC |
|
||||
| kingbase | 人大金仓 KingbaseES | KingbaseES JDBC |
|
||||
| kingbase | 人大金仓 KingbaseES | gokb Go native agent |
|
||||
| vastbase | Vastbase | Vastbase JDBC |
|
||||
| goldendb | GoldenDB | MySQL Connector/J |
|
||||
| databend | Databend | Databend JDBC |
|
||||
|
|
@ -47,16 +47,16 @@ Each agent runs as a standalone process and communicates with DBX via stdin/stdo
|
|||
|
||||
## Multi-JRE Support
|
||||
|
||||
Most Java agents target JRE 21. Native agents, such as `oracle` and `xugu`, do not require a JRE. DBX downloads and manages the JRE 21 installation automatically for Java agents.
|
||||
Most Java agents target JRE 21. Native agents, such as `oracle`, `kingbase`, and `xugu`, do not require a JRE. DBX downloads and manages the JRE 21 installation automatically for Java agents.
|
||||
|
||||
## Choosing a Driver Language
|
||||
|
||||
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 (Go/Rust)** — preferred when a usable native driver exists. See `drivers/oracle-go` (go-ora) and `drivers/xugu` as reference implementations. No JRE download or management is needed.
|
||||
- **Native (Go/Rust)** — preferred when a usable native driver exists. See `drivers/oracle-go` (go-ora), `drivers/kingbase-go` (gokb), and `drivers/xugu` 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 a native and a Java path exist for the same database, default DBX to the native one and keep the Java variant only as a compatibility fallback — see how `oracle` (go-ora native) coexists with `oracle-legacy` / `oracle-10g`.
|
||||
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`.
|
||||
|
||||
## Build
|
||||
|
||||
|
|
@ -65,10 +65,11 @@ Requires JDK 21 (Gradle toolchain auto-downloads if needed).
|
|||
```bash
|
||||
./gradlew shadowJar
|
||||
(cd drivers/oracle-go && go build -o agent .)
|
||||
(cd drivers/kingbase-go && go build -o agent .)
|
||||
(cd drivers/xugu && go build -o agent .)
|
||||
```
|
||||
|
||||
Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/oracle-go` and `drivers/xugu`.
|
||||
Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/oracle-go`, `drivers/kingbase-go`, and `drivers/xugu`.
|
||||
|
||||
### Local DBX Runtime Test
|
||||
|
||||
|
|
@ -82,7 +83,7 @@ cp agents/drivers/<db_type>/build/libs/*-all.jar ~/.dbx/agents/drivers/<db_type>
|
|||
|
||||
Restart DBX or disconnect and reconnect the database so the new agent process loads the replacement JAR.
|
||||
|
||||
Native agents such as `oracle` and `xugu` use the `agent` executable in the driver directory instead of `agent.jar`.
|
||||
Native agents such as `oracle`, `kingbase`, and `xugu` use the `agent` executable in the driver directory instead of `agent.jar`.
|
||||
|
||||
## Versioning
|
||||
|
||||
|
|
|
|||
|
|
@ -8,11 +8,11 @@ DBX 的 Agent 驱动 —— 通过 JDBC 和原生数据库驱动支持各种数
|
|||
|
||||
## 支持的数据库
|
||||
|
||||
| Agent | 数据库 | JDBC 驱动 |
|
||||
| Agent | 数据库 | 驱动 |
|
||||
|-------|----------|-------------|
|
||||
| access | Microsoft Access | UCanAccess |
|
||||
| dameng | 达梦 DM8 | DM JDBC |
|
||||
| kingbase | 人大金仓 KingbaseES | KingbaseES JDBC |
|
||||
| kingbase | 人大金仓 KingbaseES | gokb Go 原生 agent |
|
||||
| vastbase | Vastbase | Vastbase JDBC |
|
||||
| goldendb | GoldenDB | MySQL Connector/J |
|
||||
| databend | Databend | Databend JDBC |
|
||||
|
|
@ -47,16 +47,16 @@ DBX 的 Agent 驱动 —— 通过 JDBC 和原生数据库驱动支持各种数
|
|||
|
||||
## 多 JRE 支持
|
||||
|
||||
多数 Java agent 以 JRE 21 为目标。原生 agent(如 `oracle` 和 `xugu`)不需要 JRE。对 Java agent,DBX 会自动下载并管理 JRE 21 安装。
|
||||
多数 Java agent 以 JRE 21 为目标。原生 agent(如 `oracle`、`kingbase` 和 `xugu`)不需要 JRE。对 Java agent,DBX 会自动下载并管理 JRE 21 安装。
|
||||
|
||||
## 选择驱动实现语言
|
||||
|
||||
对于新 agent,只要存在成熟、许可证兼容的原生驱动,优先选择**原生(Go 或 Rust)驱动**而非 Java/JDBC agent。原生 agent 以单一自包含可执行文件发布,无需 JRE,可显著降低内存占用和启动时间 —— 完全避开 Java agent 即便空闲也要付出的 JVM 基线开销。
|
||||
|
||||
- **原生(Go/Rust)** —— 存在可用原生驱动时首选。参考 `drivers/oracle-go`(go-ora)和 `drivers/xugu`。无需 JRE 下载与管理。
|
||||
- **原生(Go/Rust)** —— 存在可用原生驱动时首选。参考 `drivers/oracle-go`(go-ora)、`drivers/kingbase-go`(gokb)和 `drivers/xugu`。无需 JRE 下载与管理。
|
||||
- **Java/JDBC** —— 当某数据库只有 JDBC 驱动,或原生驱动不成熟、缺乏维护时的默认兜底方案。多数 agent 仍属此类。
|
||||
|
||||
原生 agent 实现与 Java agent 相同的 JSON-RPC 契约和 `versions.json` 登记;它发布的是 `agent` 可执行文件而非 `agent.jar`。若同一数据库同时存在原生和 Java 路径,DBX 默认使用原生方案,仅将 Java 变体作为兼容兜底保留 —— 参见 `oracle`(go-ora 原生)与 `oracle-legacy` / `oracle-10g` 的共存方式。
|
||||
原生 agent 实现与 Java agent 相同的 JSON-RPC 契约和 `versions.json` 登记;它发布的是 `agent` 可执行文件而非 `agent.jar`。若同一数据库同时保留原生和 Java 源码实现,默认只发布原生产物;只有 Java 变体以独立兼容配置登记时才同时发布,例如 `oracle-legacy` / `oracle-10g`。
|
||||
|
||||
## 构建
|
||||
|
||||
|
|
@ -65,10 +65,11 @@ DBX 的 Agent 驱动 —— 通过 JDBC 和原生数据库驱动支持各种数
|
|||
```bash
|
||||
./gradlew shadowJar
|
||||
(cd drivers/oracle-go && go build -o agent .)
|
||||
(cd drivers/kingbase-go && go build -o agent .)
|
||||
(cd drivers/xugu && go build -o agent .)
|
||||
```
|
||||
|
||||
产物 JAR 在 `drivers/{module}/build/libs/`。原生 agent 从 `drivers/oracle-go` 和 `drivers/xugu` 构建。
|
||||
产物 JAR 在 `drivers/{module}/build/libs/`。原生 agent 从 `drivers/oracle-go`、`drivers/kingbase-go` 和 `drivers/xugu` 构建。
|
||||
|
||||
### 本地 DBX 运行时测试
|
||||
|
||||
|
|
@ -82,7 +83,7 @@ cp agents/drivers/<db_type>/build/libs/*-all.jar ~/.dbx/agents/drivers/<db_type>
|
|||
|
||||
重启 DBX 或断开重连数据库,使新 agent 进程加载替换后的 JAR。
|
||||
|
||||
`oracle` 和 `xugu` 等原生 agent 使用驱动目录下的 `agent` 可执行文件而非 `agent.jar`。
|
||||
`oracle`、`kingbase` 和 `xugu` 等原生 agent 使用驱动目录下的 `agent` 可执行文件而非 `agent.jar`。
|
||||
|
||||
## 版本管理
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,393 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
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 benchmarkResult struct {
|
||||
Server string `json:"server"`
|
||||
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"`
|
||||
StartupMS float64 `json:"startup_ms,omitempty"`
|
||||
OpenSessionsMS float64 `json:"open_sessions_ms,omitempty"`
|
||||
ReadyRSSKB int64 `json:"ready_rss_kb,omitempty"`
|
||||
OneSessionRSSKB int64 `json:"one_session_rss_kb,omitempty"`
|
||||
SessionsRSSKB int64 `json:"sessions_rss_kb,omitempty"`
|
||||
PeakRSSKB int64 `json:"peak_rss_kb,omitempty"`
|
||||
PostLoadRSSKB int64 `json:"post_load_rss_kb,omitempty"`
|
||||
}
|
||||
|
||||
type runningAgent struct {
|
||||
name string
|
||||
process *agentProcess
|
||||
startup time.Duration
|
||||
openSessions time.Duration
|
||||
readyRSSKB int64
|
||||
oneSessionRSSKB int64
|
||||
sessionsRSSKB int64
|
||||
metricsPrinted bool
|
||||
}
|
||||
|
||||
func main() {
|
||||
host := requiredEnv("KINGBASE_HOST")
|
||||
port, err := strconv.Atoi(requiredEnv("KINGBASE_PORT"))
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
database := requiredEnv("KINGBASE_DATABASE")
|
||||
username := requiredEnv("KINGBASE_USERNAME")
|
||||
password := requiredEnv("KINGBASE_PASSWORD")
|
||||
serverName := envOr("KINGBASE_SERVER", host+":"+strconv.Itoa(port))
|
||||
duration := time.Duration(envInt("BENCH_SECONDS", 4)) * time.Second
|
||||
rounds := envInt("BENCH_ROUNDS", 3)
|
||||
maxConcurrency := envInt("BENCH_MAX_CONCURRENCY", 32)
|
||||
concurrencies := envInts("BENCH_CONCURRENCIES", []int{1, 8, 32})
|
||||
benchmarkSQL := envOr("BENCH_SQL", "SELECT 1")
|
||||
benchmarkMaxRows := envInt("BENCH_MAX_ROWS", 1)
|
||||
workload := envOr("BENCH_WORKLOAD", "agent-select-literal")
|
||||
sampleMemory := os.Getenv("BENCH_SAMPLE_MEMORY") == "1"
|
||||
|
||||
agents := []struct {
|
||||
name string
|
||||
argv []string
|
||||
}{
|
||||
{name: "go-gokb", argv: []string{requiredEnv("GO_AGENT")}},
|
||||
{name: "jdbc", argv: jdbcAgentCommand(requiredEnv("JDBC_AGENT_JAR"))},
|
||||
}
|
||||
encoder := json.NewEncoder(os.Stdout)
|
||||
running := make([]*runningAgent, 0, len(agents))
|
||||
for _, candidate := range agents {
|
||||
process, startup, err := startAgent(candidate.argv)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("start %s: %w", candidate.name, err))
|
||||
}
|
||||
params := map[string]any{"host": host, "port": port, "database": database, "username": username, "password": password}
|
||||
readyRSS := readRSSKB(process.command.Process.Pid)
|
||||
var oneSessionRSS int64
|
||||
openStart := time.Now()
|
||||
for index := 0; index < maxConcurrency; index++ {
|
||||
request := cloneMap(params)
|
||||
request["agentSessionId"] = sessionID(index)
|
||||
if _, err := process.call("open_session", request); err != nil {
|
||||
panic(fmt.Errorf("open %s session %d: %w", candidate.name, index, err))
|
||||
}
|
||||
if index == 0 {
|
||||
oneSessionRSS = readRSSKB(process.command.Process.Pid)
|
||||
}
|
||||
}
|
||||
openDuration := time.Since(openStart)
|
||||
running = append(running, &runningAgent{
|
||||
name: candidate.name, process: process, startup: startup, openSessions: openDuration,
|
||||
readyRSSKB: readyRSS, oneSessionRSSKB: oneSessionRSS, sessionsRSSKB: readRSSKB(process.command.Process.Pid),
|
||||
})
|
||||
}
|
||||
defer func() {
|
||||
for _, candidate := range running {
|
||||
_ = candidate.process.close()
|
||||
}
|
||||
}()
|
||||
for _, concurrency := range concurrencies {
|
||||
if concurrency > maxConcurrency {
|
||||
continue
|
||||
}
|
||||
for round := 1; round <= rounds; round++ {
|
||||
order := running
|
||||
if round%2 == 0 {
|
||||
order = []*runningAgent{running[1], running[0]}
|
||||
}
|
||||
for _, candidate := range order {
|
||||
result := runQueryBenchmark(candidate.process, duration, concurrency, benchmarkSQL, benchmarkMaxRows, sampleMemory)
|
||||
result.Server = serverName
|
||||
result.Agent = candidate.name
|
||||
result.Workload = workload
|
||||
result.Round = round
|
||||
if !candidate.metricsPrinted {
|
||||
result.StartupMS = float64(candidate.startup.Microseconds()) / 1000
|
||||
result.OpenSessionsMS = float64(candidate.openSessions.Microseconds()) / 1000
|
||||
candidate.metricsPrinted = true
|
||||
}
|
||||
if sampleMemory {
|
||||
result.ReadyRSSKB = candidate.readyRSSKB
|
||||
result.OneSessionRSSKB = candidate.oneSessionRSSKB
|
||||
result.SessionsRSSKB = candidate.sessionsRSSKB
|
||||
result.PostLoadRSSKB = readRSSKB(candidate.process.command.Process.Pid)
|
||||
}
|
||||
if err := encoder.Encode(result); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func jdbcAgentCommand(jarPath string) []string {
|
||||
return []string{
|
||||
"java",
|
||||
"-Dfile.encoding=UTF-8",
|
||||
"-Dsun.stdout.encoding=UTF-8",
|
||||
"-Dsun.stderr.encoding=UTF-8",
|
||||
"-Djava.net.useSystemProxies=false",
|
||||
"-Dhttp.proxyHost=",
|
||||
"-Dhttps.proxyHost=",
|
||||
"-DsocksProxyHost=",
|
||||
"-Doracle.net.disableOob=true",
|
||||
"-Doracle.jdbc.javaNetNio=false",
|
||||
"-Djava.net.preferIPv4Stack=true",
|
||||
"--add-opens=java.sql/java.sql=ALL-UNNAMED",
|
||||
"-XX:TieredStopAtLevel=1",
|
||||
"-XX:+UseSerialGC",
|
||||
"-jar",
|
||||
jarPath,
|
||||
}
|
||||
}
|
||||
|
||||
func startAgent(argv []string) (*agentProcess, time.Duration, error) {
|
||||
if len(argv) == 0 {
|
||||
return nil, 0, errors.New("empty agent command")
|
||||
}
|
||||
process := &agentProcess{}
|
||||
process.command = exec.Command(argv[0], argv[1:]...)
|
||||
stdin, err := process.command.StdinPipe()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
stdout, err := process.command.StdoutPipe()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
process.command.Stderr = os.Stderr
|
||||
process.stdin = stdin
|
||||
process.reader = bufio.NewScanner(stdout)
|
||||
process.reader.Buffer(make([]byte, 64*1024), 512*1024*1024)
|
||||
start := time.Now()
|
||||
if err := process.command.Start(); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if !process.reader.Scan() || !strings.Contains(process.reader.Text(), `"ready":true`) {
|
||||
return nil, 0, fmt.Errorf("agent did not become ready: %s", process.reader.Text())
|
||||
}
|
||||
startup := time.Since(start)
|
||||
go process.readResponses()
|
||||
return process, startup, 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 {
|
||||
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(30 * 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 runQueryBenchmark(process *agentProcess, duration time.Duration, concurrency int, sqlText string, maxRows int, sampleMemory bool) benchmarkResult {
|
||||
var operations atomic.Int64
|
||||
var failures atomic.Int64
|
||||
var peakRSS atomic.Int64
|
||||
stopMemory := make(chan struct{})
|
||||
if sampleMemory {
|
||||
peakRSS.Store(readRSSKB(process.command.Process.Pid))
|
||||
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)
|
||||
params := map[string]any{"agentSessionId": sessionID(worker), "sql": sqlText, "maxRows": maxRows}
|
||||
for time.Now().Before(deadline) {
|
||||
callStart := time.Now()
|
||||
_, err := process.call("execute_query", params)
|
||||
local = append(local, float64(time.Since(callStart).Microseconds())/1000)
|
||||
operations.Add(1)
|
||||
if err != nil {
|
||||
failures.Add(1)
|
||||
}
|
||||
}
|
||||
latencies[worker] = local
|
||||
}()
|
||||
}
|
||||
workers.Wait()
|
||||
if sampleMemory {
|
||||
close(stopMemory)
|
||||
}
|
||||
elapsed := time.Since(start)
|
||||
merged := []float64{}
|
||||
for _, values := range latencies {
|
||||
merged = append(merged, values...)
|
||||
}
|
||||
sort.Float64s(merged)
|
||||
var total float64
|
||||
for _, value := range merged {
|
||||
total += value
|
||||
}
|
||||
count := operations.Load()
|
||||
return benchmarkResult{
|
||||
Concurrency: concurrency, Operations: count, Errors: failures.Load(), DurationMS: float64(elapsed.Microseconds()) / 1000,
|
||||
QPS: float64(count) / elapsed.Seconds(), MeanMS: total / float64(max(1, len(merged))),
|
||||
P50MS: percentile(merged, 0.50), P95MS: percentile(merged, 0.95), P99MS: percentile(merged, 0.99),
|
||||
PeakRSSKB: peakRSS.Load(),
|
||||
}
|
||||
}
|
||||
|
||||
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 percentile(values []float64, fraction float64) float64 {
|
||||
if len(values) == 0 {
|
||||
return 0
|
||||
}
|
||||
index := int(float64(len(values)-1) * fraction)
|
||||
return values[index]
|
||||
}
|
||||
|
||||
func sessionID(index int) string { return "bench-" + strconv.Itoa(index) }
|
||||
|
||||
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 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 envInts(name string, fallback []int) []int {
|
||||
raw := strings.TrimSpace(os.Getenv(name))
|
||||
if raw == "" {
|
||||
return fallback
|
||||
}
|
||||
result := []int{}
|
||||
for _, item := range strings.Split(raw, ",") {
|
||||
value, err := strconv.Atoi(strings.TrimSpace(item))
|
||||
if err == nil && value > 0 {
|
||||
result = append(result, value)
|
||||
}
|
||||
}
|
||||
if len(result) == 0 {
|
||||
return fallback
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
module github.com/t8y2/dbx/agents/drivers/kingbase-go
|
||||
|
||||
go 1.22
|
||||
|
||||
require gitea.com/kingbase/gokb v0.0.0-20201021123113-29bd62a876c3
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
gitea.com/kingbase/gokb v0.0.0-20201021123113-29bd62a876c3 h1:QjslQNaH5Nuap5i4nijS0OYV6GMk5kqrAmgU90zBKd4=
|
||||
gitea.com/kingbase/gokb v0.0.0-20201021123113-29bd62a876c3/go.mod h1:7lH5A1jzCXD9Nl16DzaBUOfDAT8NPrDmZwKu1p5wf94=
|
||||
|
|
@ -0,0 +1,135 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestKingbaseIntegration(t *testing.T) {
|
||||
host := os.Getenv("KINGBASE_TEST_HOST")
|
||||
portText := os.Getenv("KINGBASE_TEST_PORT")
|
||||
username := os.Getenv("KINGBASE_TEST_USERNAME")
|
||||
password := os.Getenv("KINGBASE_TEST_PASSWORD")
|
||||
if host == "" || portText == "" || username == "" || password == "" {
|
||||
t.Skip("Kingbase integration environment is not configured")
|
||||
}
|
||||
port, err := strconv.Atoi(portText)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
database := os.Getenv("KINGBASE_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:kingbase8://%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, "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)
|
||||
}
|
||||
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 sys_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)
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,882 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gitea.com/kingbase/gokb"
|
||||
)
|
||||
|
||||
const metadataTimeout = 15 * time.Second
|
||||
|
||||
var kingbaseDataTypes = []string{
|
||||
"bigint", "bigserial", "bit", "bit varying", "boolean", "bytea", "char", "character",
|
||||
"character varying", "date", "decimal", "double precision", "integer", "interval", "json",
|
||||
"jsonb", "money", "numeric", "real", "smallint", "smallserial", "serial", "text", "time",
|
||||
"time with time zone", "timestamp", "timestamp with time zone", "uuid", "varchar", "xml",
|
||||
}
|
||||
|
||||
type kingbaseMode struct {
|
||||
postgresCatalog bool
|
||||
mysqlCompat bool
|
||||
sqlServerIdentity bool
|
||||
}
|
||||
|
||||
type databaseInfo struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
type tableInfo struct {
|
||||
Name string `json:"name"`
|
||||
TableType string `json:"table_type"`
|
||||
Comment *string `json:"comment"`
|
||||
}
|
||||
|
||||
type objectInfo struct {
|
||||
Name string `json:"name"`
|
||||
ObjectType string `json:"object_type"`
|
||||
Schema string `json:"schema"`
|
||||
Comment *string `json:"comment"`
|
||||
Valid *bool `json:"valid,omitempty"`
|
||||
}
|
||||
|
||||
type metadataListConstraints struct {
|
||||
Filter string
|
||||
Limit int
|
||||
Offset int
|
||||
ObjectTypes []string
|
||||
}
|
||||
|
||||
type columnInfo struct {
|
||||
Name string `json:"name"`
|
||||
DataType string `json:"data_type"`
|
||||
IsNullable bool `json:"is_nullable"`
|
||||
ColumnDefault *string `json:"column_default"`
|
||||
IsPrimaryKey bool `json:"is_primary_key"`
|
||||
Extra *string `json:"extra"`
|
||||
Comment *string `json:"comment"`
|
||||
NumericPrecision *int `json:"numeric_precision"`
|
||||
NumericScale *int `json:"numeric_scale"`
|
||||
CharacterMaximumLength *int `json:"character_maximum_length"`
|
||||
}
|
||||
|
||||
type indexInfo struct {
|
||||
Name string `json:"name"`
|
||||
Columns []string `json:"columns"`
|
||||
IsUnique bool `json:"is_unique"`
|
||||
IsPrimary bool `json:"is_primary"`
|
||||
Filter *string `json:"filter"`
|
||||
IndexType *string `json:"index_type"`
|
||||
IncludedColumns []string `json:"included_columns"`
|
||||
Comment *string `json:"comment"`
|
||||
}
|
||||
|
||||
func (i indexInfo) MarshalJSON() ([]byte, error) {
|
||||
type alias indexInfo
|
||||
value := alias(i)
|
||||
if value.Columns == nil {
|
||||
value.Columns = []string{}
|
||||
}
|
||||
if value.IncludedColumns == nil {
|
||||
value.IncludedColumns = []string{}
|
||||
}
|
||||
return json.Marshal(value)
|
||||
}
|
||||
|
||||
type foreignKeyInfo struct {
|
||||
Name string `json:"name"`
|
||||
Column string `json:"column"`
|
||||
RefTable string `json:"ref_table"`
|
||||
RefColumn string `json:"ref_column"`
|
||||
}
|
||||
|
||||
type triggerInfo struct {
|
||||
Name string `json:"name"`
|
||||
Event string `json:"event"`
|
||||
Timing string `json:"timing"`
|
||||
}
|
||||
|
||||
func detectKingbaseMode(db *sql.DB, configuredMySQL bool) kingbaseMode {
|
||||
mode := kingbaseMode{mysqlCompat: configuredMySQL}
|
||||
if configuredMySQL {
|
||||
return mode
|
||||
}
|
||||
mode.postgresCatalog = !catalogExists(db, "sys_catalog.sys_namespace") && catalogExists(db, "pg_catalog.pg_namespace")
|
||||
if !mode.postgresCatalog {
|
||||
mode.mysqlCompat = detectMySQLCompatMode(db)
|
||||
mode.sqlServerIdentity = !mode.mysqlCompat && catalogExists(db, "sys.identity_columns")
|
||||
}
|
||||
return mode
|
||||
}
|
||||
|
||||
func catalogExists(db *sql.DB, catalog string) bool {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
rows, err := db.QueryContext(ctx, "SELECT 1 FROM "+catalog+" WHERE 1 = 0")
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return rows.Close() == nil
|
||||
}
|
||||
|
||||
func detectMySQLCompatMode(db *sql.DB) bool {
|
||||
for _, query := range []string{
|
||||
"SELECT setting FROM sys_catalog.sys_settings WHERE LOWER(name) = 'database_mode'",
|
||||
"SELECT 'mysql' FROM sys_catalog.sys_settings WHERE LOWER(name) = 'sql_mode'",
|
||||
} {
|
||||
var value string
|
||||
if db.QueryRow(query).Scan(&value) == nil && strings.EqualFold(value, "mysql") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *server) connectionInfo() (map[string]any, error) {
|
||||
db, err := s.requireDB()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var database, username, version, schema string
|
||||
err = db.QueryRow("SELECT current_database(), current_user, version(), current_schema()").Scan(&database, &username, &version, &schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return map[string]any{
|
||||
"database": database, "username": username, "version": version, "schema": schema,
|
||||
"mysql_compat_mode": s.mode.mysqlCompat,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *server) listDatabases() ([]databaseInfo, error) {
|
||||
queries := []string{
|
||||
"SELECT datname FROM sys_catalog.sys_database WHERE NOT datistemplate AND datallowconn ORDER BY datname",
|
||||
"SELECT datname FROM pg_catalog.pg_database WHERE NOT datistemplate AND datallowconn ORDER BY datname",
|
||||
"SELECT current_database()",
|
||||
}
|
||||
for _, query := range queries {
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
result := []databaseInfo{}
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if rows.Scan(&name) == nil {
|
||||
result = append(result, databaseInfo{Name: name})
|
||||
}
|
||||
}
|
||||
err = rows.Err()
|
||||
_ = rows.Close()
|
||||
if err == nil && len(result) > 0 {
|
||||
return result, nil
|
||||
}
|
||||
}
|
||||
return []databaseInfo{{Name: s.params.Database}}, nil
|
||||
}
|
||||
|
||||
func (s *server) listSchemas(visible []string) ([]string, error) {
|
||||
query := "SELECT nspname FROM sys_catalog.sys_namespace WHERE nspname NOT LIKE 'sys_temp_%' AND nspname NOT LIKE 'sys_toast_temp_%' ORDER BY nspname"
|
||||
if s.mode.postgresCatalog {
|
||||
query = "SELECT nspname FROM pg_catalog.pg_namespace WHERE nspname NOT LIKE 'pg_temp_%' AND nspname NOT LIKE 'pg_toast_temp_%' ORDER BY nspname"
|
||||
} else if s.mode.mysqlCompat {
|
||||
query = "SELECT schema_name FROM information_schema.schemata WHERE UPPER(schema_name) <> 'INFORMATION_SCHEMA' AND UPPER(schema_name) NOT LIKE 'SYS%' AND UPPER(schema_name) NOT LIKE 'XLOG%' ORDER BY schema_name"
|
||||
}
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
allowed := stringSet(visible)
|
||||
result := []string{}
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if err := rows.Scan(&name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(allowed) == 0 || allowed[strings.ToLower(name)] {
|
||||
result = append(result, name)
|
||||
}
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (s *server) listTables(schema string, constraints metadataListConstraints) ([]tableInfo, error) {
|
||||
effective, err := s.effectiveSchema(schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !constraintsAllowsTableLike(constraints) {
|
||||
return []tableInfo{}, nil
|
||||
}
|
||||
var query string
|
||||
if s.mode.mysqlCompat {
|
||||
query = "SELECT table_name, table_type, CAST(NULL AS varchar(4000)) FROM information_schema.tables WHERE table_schema = " + quoteLiteral(effective) + " ORDER BY table_name"
|
||||
} else {
|
||||
catalog := "sys_catalog"
|
||||
if s.mode.postgresCatalog {
|
||||
catalog = "pg_catalog"
|
||||
}
|
||||
query = fmt.Sprintf(`SELECT c.relname,
|
||||
CASE c.relkind WHEN 'r' THEN 'TABLE' WHEN 'p' THEN 'TABLE' WHEN 'v' THEN 'VIEW' WHEN 'm' THEN 'MATERIALIZED_VIEW' WHEN 'f' THEN 'FOREIGN_TABLE' ELSE 'TABLE' END,
|
||||
d.description
|
||||
FROM %s.%s_class c
|
||||
JOIN %s.%s_namespace n ON n.oid = c.relnamespace
|
||||
LEFT JOIN %s.%s_description d ON d.objoid = c.oid AND d.objsubid = 0
|
||||
WHERE n.nspname = %s AND c.relkind IN ('r','p','v','m','f') ORDER BY c.relname`, catalog, catalogPrefix(catalog), catalog, catalogPrefix(catalog), catalog, catalogPrefix(catalog), quoteLiteral(effective))
|
||||
}
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
result := []tableInfo{}
|
||||
for rows.Next() {
|
||||
var name, kind string
|
||||
var comment sql.NullString
|
||||
if err := rows.Scan(&name, &kind, &comment); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
item := tableInfo{Name: name, TableType: normalizeTableType(kind), Comment: nullStringPtr(comment)}
|
||||
if constraintsMatch(constraints, item.Name, item.TableType) {
|
||||
result = append(result, item)
|
||||
}
|
||||
}
|
||||
return pageTables(result, constraints), rows.Err()
|
||||
}
|
||||
|
||||
func (s *server) listObjects(schema string, constraints metadataListConstraints) ([]objectInfo, error) {
|
||||
effective, err := s.effectiveSchema(schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tables, err := s.listTables(effective, metadataListConstraints{})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make([]objectInfo, 0, len(tables))
|
||||
for _, table := range tables {
|
||||
result = append(result, objectInfo{Name: table.Name, ObjectType: table.TableType, Schema: effective, Comment: table.Comment})
|
||||
}
|
||||
if !s.mode.mysqlCompat {
|
||||
catalog := "sys_catalog"
|
||||
function := "sys"
|
||||
if s.mode.postgresCatalog {
|
||||
catalog, function = "pg_catalog", "pg"
|
||||
}
|
||||
query := fmt.Sprintf(`SELECT p.proname, CASE WHEN p.prorettype = 2278 THEN 'PROCEDURE' ELSE 'FUNCTION' END, d.description
|
||||
FROM %s.%s_proc p JOIN %s.%s_namespace n ON n.oid = p.pronamespace
|
||||
LEFT JOIN %s.%s_description d ON d.objoid = p.oid AND d.objsubid = 0
|
||||
WHERE n.nspname = %s ORDER BY p.proname`, catalog, function, catalog, function, catalog, function, quoteLiteral(effective))
|
||||
rows, queryErr := s.metadataQuery(query)
|
||||
if queryErr == nil {
|
||||
for rows.Next() {
|
||||
var name, kind string
|
||||
var comment sql.NullString
|
||||
if rows.Scan(&name, &kind, &comment) == nil {
|
||||
result = append(result, objectInfo{Name: name, ObjectType: kind, Schema: effective, Comment: nullStringPtr(comment)})
|
||||
}
|
||||
}
|
||||
_ = rows.Close()
|
||||
}
|
||||
}
|
||||
filtered := result[:0]
|
||||
for _, item := range result {
|
||||
if constraintsMatch(constraints, item.Name, item.ObjectType) {
|
||||
filtered = append(filtered, item)
|
||||
}
|
||||
}
|
||||
sort.SliceStable(filtered, func(i, j int) bool {
|
||||
if objectOrder(filtered[i].ObjectType) != objectOrder(filtered[j].ObjectType) {
|
||||
return objectOrder(filtered[i].ObjectType) < objectOrder(filtered[j].ObjectType)
|
||||
}
|
||||
return filtered[i].Name < filtered[j].Name
|
||||
})
|
||||
return pageObjects(filtered, constraints), nil
|
||||
}
|
||||
|
||||
func (s *server) completionAssistantSearch(request completionAssistantRequest) (completionAssistantResponse, error) {
|
||||
limit := request.MaxResults
|
||||
if limit <= 0 || limit > 1000 {
|
||||
limit = 100
|
||||
}
|
||||
kinds := stringSet(request.ObjectKinds)
|
||||
candidates := make([]completionAssistantCandidate, 0, limit+1)
|
||||
if kinds["column"] && request.ParentName != "" {
|
||||
schema := request.ParentSchema
|
||||
if schema == "" {
|
||||
schema = request.Schema
|
||||
}
|
||||
columns, err := s.getColumns(schema, request.ParentName)
|
||||
if err != nil {
|
||||
return completionAssistantResponse{}, err
|
||||
}
|
||||
for _, column := range columns {
|
||||
if !completionNameMatches(column.Name, request) {
|
||||
continue
|
||||
}
|
||||
dataType := column.DataType
|
||||
candidates = append(candidates, completionAssistantCandidate{
|
||||
Name: column.Name, Kind: "COLUMN", Schema: stringPtr(schema), ParentSchema: stringPtr(schema),
|
||||
ParentName: stringPtr(request.ParentName), Comment: column.Comment, DataType: &dataType,
|
||||
})
|
||||
}
|
||||
} else {
|
||||
schemas := []string{request.Schema}
|
||||
if request.GlobalSearch {
|
||||
visible, err := s.listSchemas(nil)
|
||||
if err != nil {
|
||||
return completionAssistantResponse{}, err
|
||||
}
|
||||
schemas = visible
|
||||
}
|
||||
objectTypes := request.ObjectKinds
|
||||
for _, schema := range schemas {
|
||||
objects, err := s.listObjects(schema, metadataListConstraints{ObjectTypes: objectTypes})
|
||||
if err != nil {
|
||||
return completionAssistantResponse{}, err
|
||||
}
|
||||
for _, object := range objects {
|
||||
if !completionNameMatches(object.Name, request) {
|
||||
continue
|
||||
}
|
||||
candidates = append(candidates, completionAssistantCandidate{Name: object.Name, Kind: object.ObjectType, Schema: stringPtr(object.Schema), Comment: object.Comment})
|
||||
if len(candidates) > limit {
|
||||
return completionAssistantResponse{Candidates: candidates[:limit], Incomplete: true}, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
incomplete := len(candidates) > limit
|
||||
if incomplete {
|
||||
candidates = candidates[:limit]
|
||||
}
|
||||
if candidates == nil {
|
||||
candidates = []completionAssistantCandidate{}
|
||||
}
|
||||
return completionAssistantResponse{Candidates: candidates, Incomplete: incomplete}, nil
|
||||
}
|
||||
|
||||
func completionNameMatches(name string, request completionAssistantRequest) bool {
|
||||
mask := request.Mask
|
||||
if mask == "" {
|
||||
return true
|
||||
}
|
||||
if !request.CaseSensitive {
|
||||
name = strings.ToLower(name)
|
||||
mask = strings.ToLower(mask)
|
||||
}
|
||||
if strings.EqualFold(request.MatchMode, "contains") {
|
||||
return strings.Contains(name, mask)
|
||||
}
|
||||
return strings.HasPrefix(name, mask)
|
||||
}
|
||||
|
||||
func (s *server) getColumns(schema, table string) ([]columnInfo, error) {
|
||||
effective, err := s.effectiveSchema(schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
primary, _ := s.primaryKeys(effective, table)
|
||||
if s.mode.mysqlCompat {
|
||||
return s.informationSchemaColumns(effective, table, primary)
|
||||
}
|
||||
catalog, prefix := "sys_catalog", "sys"
|
||||
if s.mode.postgresCatalog {
|
||||
catalog, prefix = "pg_catalog", "pg"
|
||||
return s.queryCatalogColumns(effective, table, primary, catalog, prefix, "pg_get_expr")
|
||||
}
|
||||
expression := "sys_get_expr"
|
||||
if s.usePgDefaultExpression {
|
||||
expression = "pg_get_expr"
|
||||
}
|
||||
result, err := s.queryCatalogColumns(effective, table, primary, catalog, prefix, expression)
|
||||
if err != nil && expression == "sys_get_expr" && isUndefinedFunction(err, expression) {
|
||||
// Some V8R6 PostgreSQL-mode databases keep sys_catalog while adbin is
|
||||
// pg_node_tree. Cache the compatible function after the exact failure.
|
||||
s.usePgDefaultExpression = true
|
||||
return s.queryCatalogColumns(effective, table, primary, catalog, prefix, "pg_get_expr")
|
||||
}
|
||||
return result, err
|
||||
}
|
||||
|
||||
func (s *server) queryCatalogColumns(
|
||||
schema, table string,
|
||||
primary map[string]bool,
|
||||
catalog, prefix, expression string,
|
||||
) ([]columnInfo, error) {
|
||||
query := fmt.Sprintf(`SELECT a.attname, format_type(a.atttypid, a.atttypmod), NOT a.attnotnull,
|
||||
%s(ad.adbin, ad.adrelid), d.description,
|
||||
CASE WHEN t.typname = 'numeric' AND a.atttypmod > 0 THEN ((a.atttypmod - 4) >> 16) & 65535 END,
|
||||
CASE WHEN t.typname = 'numeric' AND a.atttypmod > 0 THEN (a.atttypmod - 4) & 65535 END,
|
||||
CASE WHEN t.typname IN ('varchar','bpchar') AND a.atttypmod > 0 THEN a.atttypmod - 4 END
|
||||
FROM %s.%s_attribute a JOIN %s.%s_type t ON t.oid = a.atttypid
|
||||
JOIN %s.%s_class c ON c.oid = a.attrelid JOIN %s.%s_namespace n ON n.oid = c.relnamespace
|
||||
LEFT JOIN %s.%s_attrdef ad ON ad.adrelid = a.attrelid AND ad.adnum = a.attnum
|
||||
LEFT JOIN %s.%s_description d ON d.objoid = a.attrelid AND d.objsubid = a.attnum
|
||||
WHERE n.nspname = %s AND c.relname = %s AND a.attnum > 0 AND NOT a.attisdropped ORDER BY a.attnum`, expression, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, quoteLiteral(schema), quoteLiteral(table))
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
result := []columnInfo{}
|
||||
for rows.Next() {
|
||||
var name, dataType string
|
||||
var nullable bool
|
||||
var defaultValue, comment sql.NullString
|
||||
var precision, scale, length sql.NullInt64
|
||||
if err := rows.Scan(&name, &dataType, &nullable, &defaultValue, &comment, &precision, &scale, &length); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, columnInfo{Name: name, DataType: dataType, IsNullable: nullable, ColumnDefault: nullStringPtr(defaultValue), IsPrimaryKey: primary[strings.ToLower(name)], Comment: nullStringPtr(comment), NumericPrecision: nullIntPtr(precision), NumericScale: nullIntPtr(scale), CharacterMaximumLength: nullIntPtr(length)})
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if s.mode.sqlServerIdentity {
|
||||
s.applyIdentityMetadata(schema, table, result)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func isUndefinedFunction(err error, functionName string) bool {
|
||||
var driverError *gokb.Error
|
||||
undefined := errors.As(err, &driverError) && string(driverError.Code) == "42883"
|
||||
normalized := strings.ToLower(err.Error())
|
||||
undefined = undefined || strings.Contains(normalized, "does not exist") || strings.Contains(normalized, "不存在")
|
||||
return undefined && strings.Contains(normalized, strings.ToLower(functionName))
|
||||
}
|
||||
|
||||
func (s *server) informationSchemaColumns(schema, table string, primary map[string]bool) ([]columnInfo, error) {
|
||||
query := `SELECT column_name, data_type, is_nullable, column_default, numeric_precision, numeric_scale, character_maximum_length
|
||||
FROM information_schema.columns WHERE table_schema = ` + quoteLiteral(schema) + ` AND table_name = ` + quoteLiteral(table) + ` ORDER BY ordinal_position`
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
result := []columnInfo{}
|
||||
for rows.Next() {
|
||||
var name, dataType, nullable string
|
||||
var defaultValue sql.NullString
|
||||
var precision, scale, length sql.NullInt64
|
||||
if err := rows.Scan(&name, &dataType, &nullable, &defaultValue, &precision, &scale, &length); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if parsed := boundedVarcharLength(dataType); parsed != nil && !length.Valid {
|
||||
length = sql.NullInt64{Int64: int64(*parsed), Valid: true}
|
||||
}
|
||||
result = append(result, columnInfo{Name: name, DataType: dataType, IsNullable: strings.EqualFold(nullable, "YES"), ColumnDefault: nullStringPtr(defaultValue), IsPrimaryKey: primary[strings.ToLower(name)], NumericPrecision: nullIntPtr(precision), NumericScale: nullIntPtr(scale), CharacterMaximumLength: nullIntPtr(length)})
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (s *server) listIndexes(schema, table string) ([]indexInfo, error) {
|
||||
effective, err := s.effectiveSchema(schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
catalog, prefix := "sys_catalog", "sys"
|
||||
if s.mode.postgresCatalog {
|
||||
catalog, prefix = "pg_catalog", "pg"
|
||||
}
|
||||
query := fmt.Sprintf(`SELECT i.relname, am.amname, ix.indisunique, ix.indisprimary, a.attname, pos.n
|
||||
FROM %s.%s_index ix JOIN %s.%s_class t ON t.oid = ix.indrelid
|
||||
JOIN %s.%s_class i ON i.oid = ix.indexrelid JOIN %s.%s_namespace n ON n.oid = t.relnamespace
|
||||
JOIN %s.%s_am am ON am.oid = i.relam
|
||||
JOIN generate_series(1,64) pos(n) ON pos.n <= array_length(string_to_array(ix.indkey::text,' '),1)
|
||||
JOIN %s.%s_attribute a ON a.attrelid = t.oid AND a.attnum = (string_to_array(ix.indkey::text,' '))[pos.n]::int2
|
||||
WHERE n.nspname = %s AND t.relname = %s ORDER BY i.relname, pos.n`, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(table))
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
byName := map[string]*indexInfo{}
|
||||
order := []string{}
|
||||
for rows.Next() {
|
||||
var name, kind, column string
|
||||
var unique, primary bool
|
||||
var ordinal int
|
||||
if err := rows.Scan(&name, &kind, &unique, &primary, &column, &ordinal); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
item := byName[name]
|
||||
if item == nil {
|
||||
item = &indexInfo{Name: name, IsUnique: unique, IsPrimary: primary, IndexType: stringPtr(kind), Columns: []string{}, IncludedColumns: []string{}}
|
||||
byName[name] = item
|
||||
order = append(order, name)
|
||||
}
|
||||
item.Columns = append(item.Columns, column)
|
||||
}
|
||||
result := make([]indexInfo, 0, len(order))
|
||||
for _, name := range order {
|
||||
result = append(result, *byName[name])
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (s *server) listForeignKeys(schema, table string) ([]foreignKeyInfo, error) {
|
||||
effective, err := s.effectiveSchema(schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
query := `SELECT fk.constraint_name, fk.column_name, pk.table_name, pk.column_name
|
||||
FROM information_schema.table_constraints tc
|
||||
JOIN information_schema.key_column_usage fk ON fk.constraint_schema = tc.constraint_schema AND fk.constraint_name = tc.constraint_name AND fk.table_schema = tc.table_schema AND fk.table_name = tc.table_name
|
||||
JOIN information_schema.referential_constraints rc ON rc.constraint_schema = tc.constraint_schema AND rc.constraint_name = tc.constraint_name
|
||||
JOIN information_schema.key_column_usage pk ON pk.constraint_schema = rc.unique_constraint_schema AND pk.constraint_name = rc.unique_constraint_name AND pk.ordinal_position = fk.position_in_unique_constraint
|
||||
WHERE tc.table_schema = ` + quoteLiteral(effective) + ` AND tc.table_name = ` + quoteLiteral(table) + ` AND tc.constraint_type = 'FOREIGN KEY' ORDER BY fk.constraint_name, fk.ordinal_position`
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
result := []foreignKeyInfo{}
|
||||
for rows.Next() {
|
||||
var item foreignKeyInfo
|
||||
if err := rows.Scan(&item.Name, &item.Column, &item.RefTable, &item.RefColumn); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (s *server) listTriggers(schema, table string) ([]triggerInfo, error) {
|
||||
effective, err := s.effectiveSchema(schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
catalog, prefix := "sys_catalog", "sys"
|
||||
if s.mode.postgresCatalog {
|
||||
catalog, prefix = "pg_catalog", "pg"
|
||||
}
|
||||
query := fmt.Sprintf(`SELECT tg.tgname,
|
||||
trim(trailing ',' FROM (CASE WHEN (tg.tgtype & 4) <> 0 THEN 'INSERT,' ELSE '' END || CASE WHEN (tg.tgtype & 8) <> 0 THEN 'DELETE,' ELSE '' END || CASE WHEN (tg.tgtype & 16) <> 0 THEN 'UPDATE,' ELSE '' END || CASE WHEN (tg.tgtype & 32) <> 0 THEN 'TRUNCATE,' ELSE '' END)), tg.tgtype
|
||||
FROM %s.%s_trigger tg JOIN %s.%s_class c ON c.oid = tg.tgrelid JOIN %s.%s_namespace n ON n.oid = c.relnamespace
|
||||
WHERE n.nspname = %s AND c.relname = %s AND NOT tg.tgisinternal ORDER BY tg.tgname`, catalog, prefix, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(table))
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
result := []triggerInfo{}
|
||||
for rows.Next() {
|
||||
var name, event string
|
||||
var triggerType int
|
||||
if err := rows.Scan(&name, &event, &triggerType); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, triggerInfo{Name: name, Event: event, Timing: decodeTriggerTiming(triggerType)})
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (s *server) getObjectSource(schema, name, objectType string) (map[string]any, error) {
|
||||
effective, err := s.effectiveSchema(schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
source := ""
|
||||
kind := strings.ToUpper(objectType)
|
||||
if kind == "VIEW" || kind == "MATERIALIZED_VIEW" {
|
||||
if s.mode.mysqlCompat {
|
||||
err = s.requireDBQueryRow("SELECT view_definition FROM information_schema.views WHERE table_schema = "+quoteLiteral(effective)+" AND table_name = "+quoteLiteral(name), &source)
|
||||
} else {
|
||||
catalog, prefix, function := "sys_catalog", "sys", "sys_get_viewdef"
|
||||
if s.mode.postgresCatalog {
|
||||
catalog, prefix, function = "pg_catalog", "pg", "pg_get_viewdef"
|
||||
}
|
||||
query := fmt.Sprintf("SELECT %s(c.oid) FROM %s.%s_class c JOIN %s.%s_namespace n ON n.oid=c.relnamespace WHERE n.nspname=%s AND c.relname=%s LIMIT 1", function, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(name))
|
||||
err = s.requireDBQueryRow(query, &source)
|
||||
}
|
||||
} else if kind == "FUNCTION" || kind == "PROCEDURE" {
|
||||
catalog, prefix, function := "sys_catalog", "sys", "sys_get_functiondef"
|
||||
if s.mode.postgresCatalog {
|
||||
catalog, prefix, function = "pg_catalog", "pg", "pg_get_functiondef"
|
||||
}
|
||||
query := fmt.Sprintf("SELECT %s(p.oid) FROM %s.%s_proc p JOIN %s.%s_namespace n ON n.oid=p.pronamespace WHERE n.nspname=%s AND p.proname=%s ORDER BY CASE WHEN p.prorettype=2278 THEN 0 ELSE 1 END LIMIT 1", function, catalog, prefix, catalog, prefix, quoteLiteral(effective), quoteLiteral(name))
|
||||
err = s.requireDBQueryRow(query, &source)
|
||||
}
|
||||
if err != nil && err != sql.ErrNoRows {
|
||||
return nil, err
|
||||
}
|
||||
return map[string]any{"name": name, "object_type": objectType, "schema": effective, "source": source}, nil
|
||||
}
|
||||
|
||||
func (s *server) getTableDDL(schema, table string) (string, error) {
|
||||
effective, err := s.effectiveSchema(schema)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
columns, err := s.getColumns(effective, table)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
definitions := make([]string, 0, len(columns)+1)
|
||||
primary := []string{}
|
||||
for _, column := range columns {
|
||||
definition := quoteIdentifier(column.Name) + " " + column.DataType
|
||||
if !column.IsNullable {
|
||||
definition += " NOT NULL"
|
||||
}
|
||||
if column.ColumnDefault != nil && *column.ColumnDefault != "" {
|
||||
definition += " DEFAULT " + *column.ColumnDefault
|
||||
}
|
||||
definitions = append(definitions, definition)
|
||||
if column.IsPrimaryKey {
|
||||
primary = append(primary, quoteIdentifier(column.Name))
|
||||
}
|
||||
}
|
||||
if len(primary) > 0 {
|
||||
definitions = append(definitions, "PRIMARY KEY ("+strings.Join(primary, ", ")+")")
|
||||
}
|
||||
return "CREATE TABLE " + quoteIdentifier(effective) + "." + quoteIdentifier(table) + " (\n " + strings.Join(definitions, ",\n ") + "\n);", nil
|
||||
}
|
||||
|
||||
func (s *server) getExplainInfo(sqlText string) (string, error) {
|
||||
rows, err := s.metadataQuery("EXPLAIN " + trimStatementSQL(sqlText))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer rows.Close()
|
||||
lines := []string{}
|
||||
for rows.Next() {
|
||||
var line string
|
||||
if err := rows.Scan(&line); err != nil {
|
||||
return "", err
|
||||
}
|
||||
lines = append(lines, line)
|
||||
}
|
||||
return strings.Join(lines, "\n"), rows.Err()
|
||||
}
|
||||
|
||||
func (s *server) metadataQuery(query string) (*sql.Rows, error) {
|
||||
db, err := s.requireDB()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// These are bounded, internally generated statements. Calling Query without
|
||||
// arguments keeps gokb on its single-round-trip simple-query path.
|
||||
return db.Query(query)
|
||||
}
|
||||
|
||||
func (s *server) requireDBQueryRow(query string, destination ...any) error {
|
||||
db, err := s.requireDB()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), metadataTimeout)
|
||||
defer cancel()
|
||||
return db.QueryRowContext(ctx, query).Scan(destination...)
|
||||
}
|
||||
|
||||
func (s *server) effectiveSchema(schema string) (string, error) {
|
||||
if strings.TrimSpace(schema) != "" {
|
||||
return strings.TrimSpace(schema), nil
|
||||
}
|
||||
var current sql.NullString
|
||||
if err := s.requireDBQueryRow("SELECT current_schema()", ¤t); err == nil && current.Valid && current.String != "" {
|
||||
return current.String, nil
|
||||
}
|
||||
if s.params.Username != "" {
|
||||
return s.params.Username, nil
|
||||
}
|
||||
return "public", nil
|
||||
}
|
||||
|
||||
func (s *server) primaryKeys(schema, table string) (map[string]bool, error) {
|
||||
query := `SELECT kcu.column_name FROM information_schema.table_constraints tc
|
||||
JOIN information_schema.key_column_usage kcu ON kcu.constraint_schema=tc.constraint_schema AND kcu.constraint_name=tc.constraint_name AND kcu.table_schema=tc.table_schema AND kcu.table_name=tc.table_name
|
||||
WHERE tc.table_schema=` + quoteLiteral(schema) + ` AND tc.table_name=` + quoteLiteral(table) + ` AND tc.constraint_type='PRIMARY KEY' ORDER BY kcu.ordinal_position`
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
result := map[string]bool{}
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if err := rows.Scan(&name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result[strings.ToLower(name)] = true
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (s *server) applyIdentityMetadata(schema, table string, columns []columnInfo) {
|
||||
query := `SELECT a.attname, ic.seed_value, ic.increment_value FROM sys.identity_columns ic
|
||||
JOIN sys_catalog.sys_class c ON c.oid=ic.object_id JOIN sys_catalog.sys_namespace n ON n.oid=c.relnamespace
|
||||
JOIN sys_catalog.sys_attribute a ON a.attrelid=c.oid AND a.attnum=ic.column_id
|
||||
WHERE n.nspname=` + quoteLiteral(schema) + ` AND c.relname=` + quoteLiteral(table)
|
||||
rows, err := s.metadataQuery(query)
|
||||
if err != nil {
|
||||
s.mode.sqlServerIdentity = false
|
||||
return
|
||||
}
|
||||
defer rows.Close()
|
||||
byName := map[string]*columnInfo{}
|
||||
for i := range columns {
|
||||
byName[strings.ToLower(columns[i].Name)] = &columns[i]
|
||||
}
|
||||
for rows.Next() {
|
||||
var name string
|
||||
var seed, increment sql.NullString
|
||||
if rows.Scan(&name, &seed, &increment) == nil {
|
||||
if column := byName[strings.ToLower(name)]; column != nil {
|
||||
extra := "IDENTITY"
|
||||
if seed.Valid && increment.Valid {
|
||||
extra = "IDENTITY(" + seed.String + "," + increment.String + ")"
|
||||
}
|
||||
column.Extra = &extra
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func catalogPrefix(catalog string) string {
|
||||
if catalog == "pg_catalog" {
|
||||
return "pg"
|
||||
}
|
||||
return "sys"
|
||||
}
|
||||
|
||||
func normalizeTableType(value string) string {
|
||||
normalized := strings.ToUpper(strings.ReplaceAll(strings.TrimSpace(value), " ", "_"))
|
||||
switch normalized {
|
||||
case "BASE_TABLE", "PARTITIONED_TABLE":
|
||||
return "TABLE"
|
||||
case "MATERIALIZED_VIEW", "FOREIGN_TABLE", "VIEW", "TABLE":
|
||||
return normalized
|
||||
default:
|
||||
return "TABLE"
|
||||
}
|
||||
}
|
||||
|
||||
func decodeTriggerTiming(triggerType int) string {
|
||||
if triggerType&(1<<6) != 0 {
|
||||
return "INSTEAD OF"
|
||||
}
|
||||
if triggerType&(1<<1) != 0 {
|
||||
return "BEFORE"
|
||||
}
|
||||
return "AFTER"
|
||||
}
|
||||
|
||||
func boundedVarcharLength(dataType string) *int {
|
||||
lower := strings.ToLower(strings.TrimSpace(dataType))
|
||||
for _, prefix := range []string{"varchar", "character varying"} {
|
||||
if strings.HasPrefix(lower, prefix) {
|
||||
value := strings.TrimSpace(strings.TrimSuffix(strings.TrimPrefix(lower, prefix), ")"))
|
||||
value = strings.TrimPrefix(value, "(")
|
||||
if number, err := strconv.Atoi(strings.TrimSpace(value)); err == nil && number >= 0 {
|
||||
return &number
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func constraintsAllowsTableLike(constraints metadataListConstraints) bool {
|
||||
if len(constraints.ObjectTypes) == 0 {
|
||||
return true
|
||||
}
|
||||
for _, kind := range constraints.ObjectTypes {
|
||||
switch normalizeTableType(kind) {
|
||||
case "TABLE", "VIEW", "MATERIALIZED_VIEW", "FOREIGN_TABLE":
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func constraintsMatch(constraints metadataListConstraints, name, kind string) bool {
|
||||
if filter := strings.TrimSpace(constraints.Filter); filter != "" && !strings.Contains(strings.ToLower(name), strings.ToLower(filter)) {
|
||||
return false
|
||||
}
|
||||
if len(constraints.ObjectTypes) == 0 {
|
||||
return true
|
||||
}
|
||||
for _, allowed := range constraints.ObjectTypes {
|
||||
if strings.EqualFold(normalizeTableType(allowed), normalizeTableType(kind)) || strings.EqualFold(allowed, kind) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func pageTables(items []tableInfo, constraints metadataListConstraints) []tableInfo {
|
||||
start, end := pageBounds(len(items), constraints.Offset, constraints.Limit)
|
||||
return items[start:end]
|
||||
}
|
||||
|
||||
func pageObjects(items []objectInfo, constraints metadataListConstraints) []objectInfo {
|
||||
start, end := pageBounds(len(items), constraints.Offset, constraints.Limit)
|
||||
return items[start:end]
|
||||
}
|
||||
|
||||
func pageBounds(length, offset, limit int) (int, int) {
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
if offset > length {
|
||||
offset = length
|
||||
}
|
||||
end := length
|
||||
if limit > 0 && offset+limit < end {
|
||||
end = offset + limit
|
||||
}
|
||||
return offset, end
|
||||
}
|
||||
|
||||
func objectOrder(kind string) int {
|
||||
switch strings.ToUpper(kind) {
|
||||
case "TABLE":
|
||||
return 0
|
||||
case "VIEW":
|
||||
return 1
|
||||
case "MATERIALIZED_VIEW":
|
||||
return 2
|
||||
case "FOREIGN_TABLE":
|
||||
return 3
|
||||
case "PROCEDURE":
|
||||
return 4
|
||||
case "FUNCTION":
|
||||
return 5
|
||||
default:
|
||||
return 9
|
||||
}
|
||||
}
|
||||
|
||||
func stringSet(values []string) map[string]bool {
|
||||
result := map[string]bool{}
|
||||
for _, value := range values {
|
||||
result[strings.ToLower(value)] = true
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func nullStringPtr(value sql.NullString) *string {
|
||||
if !value.Valid {
|
||||
return nil
|
||||
}
|
||||
return &value.String
|
||||
}
|
||||
|
||||
func nullIntPtr(value sql.NullInt64) *int {
|
||||
if !value.Valid {
|
||||
return nil
|
||||
}
|
||||
converted := int(value.Int64)
|
||||
return &converted
|
||||
}
|
||||
|
|
@ -0,0 +1,985 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
_ "gitea.com/kingbase/gokb"
|
||||
)
|
||||
|
||||
const (
|
||||
protocolVersion = 2
|
||||
defaultMaxRows = 10000
|
||||
legacyAgentSessionID = "__legacy__"
|
||||
maxAgentSessions = 256
|
||||
defaultConnectTimeout = 15 * time.Second
|
||||
)
|
||||
|
||||
type request struct {
|
||||
ID json.RawMessage `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params map[string]json.RawMessage `json:"params"`
|
||||
}
|
||||
|
||||
type response struct {
|
||||
JSONRPC string `json:"jsonrpc,omitempty"`
|
||||
ID json.RawMessage `json:"id,omitempty"`
|
||||
Result any `json:"result,omitempty"`
|
||||
Error *rpcError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type rpcError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
type connectParams struct {
|
||||
Host string `json:"host"`
|
||||
Port int `json:"port"`
|
||||
Database string `json:"database"`
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
URLParams string `json:"url_params"`
|
||||
ConnectionString string `json:"connection_string"`
|
||||
MySQLCompatMode bool `json:"mysql_compat_mode"`
|
||||
SSL bool `json:"ssl"`
|
||||
CACertPath string `json:"ca_cert_path"`
|
||||
ClientCertPath string `json:"client_cert_path"`
|
||||
ClientKeyPath string `json:"client_key_path"`
|
||||
}
|
||||
|
||||
type queryOptions struct {
|
||||
SQL string `json:"sql"`
|
||||
Database string `json:"database"`
|
||||
Schema string `json:"schema"`
|
||||
MaxRows int `json:"maxRows"`
|
||||
FetchSize int `json:"fetchSize"`
|
||||
TimeoutSecs int `json:"timeoutSecs"`
|
||||
}
|
||||
|
||||
type completionAssistantRequest struct {
|
||||
ConnectionID string `json:"connection_id"`
|
||||
Database string `json:"database"`
|
||||
Schema string `json:"schema"`
|
||||
ObjectKinds []string `json:"object_kinds"`
|
||||
Mask string `json:"mask"`
|
||||
CaseSensitive bool `json:"case_sensitive"`
|
||||
GlobalSearch bool `json:"global_search"`
|
||||
MaxResults int `json:"max_results"`
|
||||
ParentSchema string `json:"parent_schema"`
|
||||
ParentName string `json:"parent_name"`
|
||||
MatchMode string `json:"match_mode"`
|
||||
}
|
||||
|
||||
type completionAssistantCandidate struct {
|
||||
Name string `json:"name"`
|
||||
Kind string `json:"kind"`
|
||||
Database *string `json:"database"`
|
||||
Schema *string `json:"schema"`
|
||||
ParentSchema *string `json:"parent_schema"`
|
||||
ParentName *string `json:"parent_name"`
|
||||
Comment *string `json:"comment"`
|
||||
DataType *string `json:"data_type"`
|
||||
}
|
||||
|
||||
type completionAssistantResponse struct {
|
||||
Candidates []completionAssistantCandidate `json:"candidates"`
|
||||
Incomplete bool `json:"incomplete"`
|
||||
FallbackUsed bool `json:"fallback_used"`
|
||||
}
|
||||
|
||||
type queryResult struct {
|
||||
Columns []string `json:"columns"`
|
||||
ColumnTypes []string `json:"column_types"`
|
||||
Rows [][]any `json:"rows"`
|
||||
AffectedRows int64 `json:"affected_rows"`
|
||||
ExecutionTimeMS int64 `json:"execution_time_ms"`
|
||||
Truncated bool `json:"truncated"`
|
||||
}
|
||||
|
||||
type queryPageResult struct {
|
||||
Columns []string `json:"columns"`
|
||||
ColumnTypes []string `json:"column_types"`
|
||||
Rows [][]any `json:"rows"`
|
||||
AffectedRows int64 `json:"affected_rows"`
|
||||
ExecutionTimeMS int64 `json:"execution_time_ms"`
|
||||
Truncated bool `json:"truncated"`
|
||||
SessionID *string `json:"session_id"`
|
||||
HasMore bool `json:"has_more"`
|
||||
}
|
||||
|
||||
type querySession struct {
|
||||
rows *sql.Rows
|
||||
columns []string
|
||||
columnTypes []string
|
||||
pending []any
|
||||
remaining int
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
type server struct {
|
||||
db *sql.DB
|
||||
params connectParams
|
||||
mode kingbaseMode
|
||||
usePgDefaultExpression bool
|
||||
currentSchema string
|
||||
schemaSet bool
|
||||
sessions map[string]*querySession
|
||||
nextSessionID uint64
|
||||
activeCancelMu sync.Mutex
|
||||
activeCancel context.CancelFunc
|
||||
}
|
||||
|
||||
type agentSession struct {
|
||||
server *server
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
type runtimeServer struct {
|
||||
mu sync.RWMutex
|
||||
sessions map[string]*agentSession
|
||||
}
|
||||
|
||||
func main() {
|
||||
runtime := &runtimeServer{sessions: map[string]*agentSession{}}
|
||||
encoder := json.NewEncoder(os.Stdout)
|
||||
var encoderMu sync.Mutex
|
||||
var requests sync.WaitGroup
|
||||
fmt.Fprintln(os.Stdout, `{"ready":true}`)
|
||||
|
||||
scanner := bufio.NewScanner(os.Stdin)
|
||||
scanner.Buffer(make([]byte, 0, 64*1024), 512*1024*1024)
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
var envelope request
|
||||
if json.Unmarshal([]byte(line), &envelope) == nil && envelope.Method == "shutdown" {
|
||||
requests.Wait()
|
||||
resp, _ := runtime.handleLine(line)
|
||||
encoderMu.Lock()
|
||||
_ = encoder.Encode(resp)
|
||||
encoderMu.Unlock()
|
||||
return
|
||||
}
|
||||
requests.Add(1)
|
||||
go func(line string) {
|
||||
defer requests.Done()
|
||||
resp, _ := runtime.handleLine(line)
|
||||
encoderMu.Lock()
|
||||
defer encoderMu.Unlock()
|
||||
if err := encoder.Encode(resp); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "failed to write response: %v\n", err)
|
||||
}
|
||||
}(line)
|
||||
}
|
||||
requests.Wait()
|
||||
}
|
||||
|
||||
func (r *runtimeServer) handleLine(line string) (response, bool) {
|
||||
var req request
|
||||
if err := json.Unmarshal([]byte(line), &req); err != nil {
|
||||
return errorResponse(nil, err), false
|
||||
}
|
||||
if len(req.ID) == 0 {
|
||||
req.ID = json.RawMessage("1")
|
||||
}
|
||||
result, shutdown, err := r.dispatch(req.Method, req.Params)
|
||||
if err != nil {
|
||||
return errorResponse(req.ID, err), false
|
||||
}
|
||||
return response{JSONRPC: "2.0", ID: req.ID, Result: result}, shutdown
|
||||
}
|
||||
|
||||
func (r *runtimeServer) dispatch(method string, params map[string]json.RawMessage) (any, bool, error) {
|
||||
switch method {
|
||||
case "handshake":
|
||||
return map[string]any{
|
||||
"protocolVersion": protocolVersion,
|
||||
"agentProtocolVersion": protocolVersion,
|
||||
"capabilities": []string{
|
||||
"connect", "test_connection", "metadata", "query", "paged_query", "transaction", "ddl", "multi_session",
|
||||
},
|
||||
}, false, nil
|
||||
case "open_session":
|
||||
id := stringParam(params, "agentSessionId")
|
||||
if id == "" {
|
||||
return nil, false, errors.New("agentSessionId is required")
|
||||
}
|
||||
var cp connectParams
|
||||
if err := decodeParams(params, &cp); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
return map[string]bool{"ok": true}, false, r.openSession(id, cp)
|
||||
case "close_session":
|
||||
return map[string]bool{"ok": true}, false, r.closeSession(stringParam(params, "agentSessionId"))
|
||||
case "validate_session":
|
||||
session, err := r.session(stringParam(params, "agentSessionId"))
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
session.mu.Lock()
|
||||
defer session.mu.Unlock()
|
||||
return map[string]bool{"ok": true}, false, session.server.validateConnection()
|
||||
case "cancel_session":
|
||||
session, err := r.session(stringParam(params, "agentSessionId"))
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
session.server.cancelActiveQuery()
|
||||
return map[string]bool{"ok": true}, false, nil
|
||||
case "test_connection":
|
||||
return newServer().dispatch(method, params)
|
||||
case "connect":
|
||||
var cp connectParams
|
||||
if err := decodeParams(params, &cp); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
_ = r.closeSession(legacyAgentSessionID)
|
||||
return map[string]bool{"ok": true}, false, r.openSession(legacyAgentSessionID, cp)
|
||||
case "disconnect":
|
||||
return map[string]bool{"ok": true}, false, r.closeSession(legacyAgentSessionID)
|
||||
case "shutdown":
|
||||
return map[string]bool{"ok": true}, true, r.closeAllSessions()
|
||||
default:
|
||||
id := stringParam(params, "agentSessionId")
|
||||
if id == "" {
|
||||
id = legacyAgentSessionID
|
||||
}
|
||||
session, err := r.session(id)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
session.mu.Lock()
|
||||
defer session.mu.Unlock()
|
||||
return session.server.dispatch(method, params)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *runtimeServer) openSession(id string, cp connectParams) error {
|
||||
r.mu.Lock()
|
||||
if _, exists := r.sessions[id]; exists {
|
||||
r.mu.Unlock()
|
||||
return fmt.Errorf("agent session already exists: %s", id)
|
||||
}
|
||||
if len(r.sessions) >= maxAgentSessions {
|
||||
r.mu.Unlock()
|
||||
return fmt.Errorf("agent session limit reached: %d", maxAgentSessions)
|
||||
}
|
||||
r.mu.Unlock()
|
||||
|
||||
s := newServer()
|
||||
if err := s.connect(cp); err != nil {
|
||||
return err
|
||||
}
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if _, exists := r.sessions[id]; exists {
|
||||
_ = s.disconnect()
|
||||
return fmt.Errorf("agent session already exists: %s", id)
|
||||
}
|
||||
r.sessions[id] = &agentSession{server: s}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *runtimeServer) session(id string) (*agentSession, error) {
|
||||
r.mu.RLock()
|
||||
session := r.sessions[id]
|
||||
r.mu.RUnlock()
|
||||
if session == nil {
|
||||
return nil, fmt.Errorf("agent session not found: %s", id)
|
||||
}
|
||||
return session, nil
|
||||
}
|
||||
|
||||
func (r *runtimeServer) closeSession(id string) error {
|
||||
r.mu.Lock()
|
||||
session := r.sessions[id]
|
||||
delete(r.sessions, id)
|
||||
r.mu.Unlock()
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
session.server.cancelActiveQuery()
|
||||
session.mu.Lock()
|
||||
defer session.mu.Unlock()
|
||||
return session.server.disconnect()
|
||||
}
|
||||
|
||||
func (r *runtimeServer) closeAllSessions() error {
|
||||
r.mu.RLock()
|
||||
ids := make([]string, 0, len(r.sessions))
|
||||
for id := range r.sessions {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
r.mu.RUnlock()
|
||||
var firstErr error
|
||||
for _, id := range ids {
|
||||
if err := r.closeSession(id); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
return firstErr
|
||||
}
|
||||
|
||||
func newServer() *server {
|
||||
return &server{sessions: map[string]*querySession{}}
|
||||
}
|
||||
|
||||
func (s *server) dispatch(method string, params map[string]json.RawMessage) (any, bool, error) {
|
||||
switch method {
|
||||
case "handshake":
|
||||
return map[string]any{
|
||||
"protocolVersion": protocolVersion,
|
||||
"agentProtocolVersion": protocolVersion,
|
||||
"capabilities": []string{"connect", "test_connection", "metadata", "query", "paged_query", "transaction", "ddl"},
|
||||
}, false, nil
|
||||
case "connect":
|
||||
var cp connectParams
|
||||
if err := decodeParams(params, &cp); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
return map[string]bool{"ok": true}, false, s.connect(cp)
|
||||
case "test_connection":
|
||||
var cp connectParams
|
||||
if err := decodeParams(params, &cp); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
db, err := openDB(cp)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
defer db.Close()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), defaultConnectTimeout)
|
||||
defer cancel()
|
||||
return map[string]bool{"ok": true}, false, db.PingContext(ctx)
|
||||
case "validate_connection":
|
||||
return map[string]bool{"ok": true}, false, s.validateConnection()
|
||||
case "connection_info":
|
||||
info, err := s.connectionInfo()
|
||||
return info, false, err
|
||||
case "list_databases":
|
||||
result, err := s.listDatabases()
|
||||
return result, false, err
|
||||
case "list_schemas":
|
||||
result, err := s.listSchemas(stringSliceParam(params, "visible_schemas"))
|
||||
return result, false, err
|
||||
case "list_tables":
|
||||
result, err := s.listTables(stringParam(params, "schema"), metadataListConstraintsFromParams(params))
|
||||
return result, false, err
|
||||
case "list_objects":
|
||||
result, err := s.listObjects(stringParam(params, "schema"), metadataListConstraintsFromParams(params))
|
||||
return result, false, err
|
||||
case "list_data_types":
|
||||
return kingbaseDataTypes, false, nil
|
||||
case "completion_assistant_search_v1":
|
||||
var request completionAssistantRequest
|
||||
if err := decodeParams(params, &request); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
result, err := s.completionAssistantSearch(request)
|
||||
return result, false, err
|
||||
case "get_columns":
|
||||
result, err := s.getColumns(stringParam(params, "schema"), stringParam(params, "table"))
|
||||
return result, false, err
|
||||
case "list_indexes":
|
||||
result, err := s.listIndexes(stringParam(params, "schema"), stringParam(params, "table"))
|
||||
return result, false, err
|
||||
case "list_foreign_keys":
|
||||
result, err := s.listForeignKeys(stringParam(params, "schema"), stringParam(params, "table"))
|
||||
return result, false, err
|
||||
case "list_triggers":
|
||||
result, err := s.listTriggers(stringParam(params, "schema"), stringParam(params, "table"))
|
||||
return result, false, err
|
||||
case "get_object_source":
|
||||
result, err := s.getObjectSource(stringParam(params, "schema"), stringParam(params, "name"), stringParam(params, "object_type"))
|
||||
return result, false, err
|
||||
case "get_table_ddl":
|
||||
result, err := s.getTableDDL(stringParam(params, "schema"), stringParam(params, "table"))
|
||||
return result, false, err
|
||||
case "get_explain_info":
|
||||
result, err := s.getExplainInfo(stringParam(params, "sql"))
|
||||
return map[string]any{"plan": result, "has_actual_stats": false}, false, err
|
||||
case "execute_query":
|
||||
var opts queryOptions
|
||||
if err := decodeParams(params, &opts); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
result, err := s.executeQuery(opts)
|
||||
return result, false, err
|
||||
case "execute_query_page", "start_table_read":
|
||||
var opts queryOptions
|
||||
if err := decodeParams(params, &opts); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
result, err := s.executeQueryPage(opts, intParam(params, "pageSize"))
|
||||
return result, false, err
|
||||
case "fetch_query_page", "fetch_table_read_page":
|
||||
result, err := s.fetchQueryPage(stringParam(params, "sessionId"), intParam(params, "pageSize"))
|
||||
return result, false, err
|
||||
case "close_query_session", "close_table_read_session":
|
||||
return s.closeQuerySession(stringParam(params, "sessionId")), false, nil
|
||||
case "execute_transaction":
|
||||
result, err := s.executeTransaction(params)
|
||||
return result, false, err
|
||||
case "execute_batch":
|
||||
result, err := s.executeBatch(params)
|
||||
return result, false, err
|
||||
case "disconnect":
|
||||
return map[string]bool{"ok": true}, false, s.disconnect()
|
||||
case "shutdown":
|
||||
return map[string]bool{"ok": true}, true, s.disconnect()
|
||||
default:
|
||||
return nil, false, fmt.Errorf("unknown method: %s", method)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *server) connect(cp connectParams) error {
|
||||
_ = s.disconnect()
|
||||
db, err := openDB(cp)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), defaultConnectTimeout)
|
||||
defer cancel()
|
||||
if err := db.PingContext(ctx); err != nil {
|
||||
_ = db.Close()
|
||||
return err
|
||||
}
|
||||
s.db = db
|
||||
s.params = cp
|
||||
s.mode = detectKingbaseMode(db, cp.MySQLCompatMode)
|
||||
s.usePgDefaultExpression = false
|
||||
return nil
|
||||
}
|
||||
|
||||
func openDB(cp connectParams) (*sql.DB, error) {
|
||||
dsn := buildDSN(cp)
|
||||
db, err := sql.Open("kingbase", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Each protocol session is serialized and owns one database connection.
|
||||
// Keeping a single physical connection preserves session state such as
|
||||
// search_path and avoids extra pool coordination on the hot query path.
|
||||
db.SetMaxOpenConns(1)
|
||||
db.SetMaxIdleConns(1)
|
||||
db.SetConnMaxLifetime(5 * time.Minute)
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func (s *server) disconnect() error {
|
||||
s.cancelActiveQuery()
|
||||
s.closeAllQuerySessions()
|
||||
if s.db == nil {
|
||||
return nil
|
||||
}
|
||||
err := s.db.Close()
|
||||
s.db = nil
|
||||
s.usePgDefaultExpression = false
|
||||
s.currentSchema = ""
|
||||
s.schemaSet = false
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *server) validateConnection() error {
|
||||
db, err := s.requireDB()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
return db.PingContext(ctx)
|
||||
}
|
||||
|
||||
func (s *server) requireDB() (*sql.DB, error) {
|
||||
if s.db == nil {
|
||||
return nil, errors.New("not connected")
|
||||
}
|
||||
return s.db, nil
|
||||
}
|
||||
|
||||
func (s *server) beginOperation(timeoutSecs int) (context.Context, context.CancelFunc) {
|
||||
ctx := context.Background()
|
||||
var cancel context.CancelFunc
|
||||
if timeoutSecs > 0 {
|
||||
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeoutSecs)*time.Second)
|
||||
} else {
|
||||
ctx, cancel = context.WithCancel(ctx)
|
||||
}
|
||||
s.activeCancelMu.Lock()
|
||||
s.activeCancel = cancel
|
||||
s.activeCancelMu.Unlock()
|
||||
return ctx, cancel
|
||||
}
|
||||
|
||||
func (s *server) endOperation(cancel context.CancelFunc) {
|
||||
cancel()
|
||||
s.activeCancelMu.Lock()
|
||||
s.activeCancel = nil
|
||||
s.activeCancelMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *server) cancelActiveQuery() {
|
||||
s.activeCancelMu.Lock()
|
||||
cancel := s.activeCancel
|
||||
s.activeCancelMu.Unlock()
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *server) executeQuery(opts queryOptions) (queryResult, error) {
|
||||
start := time.Now()
|
||||
if err := s.setSchema(opts.Schema); err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
sqlText := trimStatementSQL(opts.SQL)
|
||||
if isQuerySQL(sqlText) {
|
||||
rows, cancel, err := s.queryRows(sqlText, opts.TimeoutSecs)
|
||||
if err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
defer func() {
|
||||
_ = rows.Close()
|
||||
s.endOperation(cancel)
|
||||
}()
|
||||
maxRows := opts.MaxRows
|
||||
if maxRows <= 0 {
|
||||
maxRows = defaultMaxRows
|
||||
}
|
||||
result, err := readRows(rows, maxRows)
|
||||
result.ExecutionTimeMS = time.Since(start).Milliseconds()
|
||||
return result, err
|
||||
}
|
||||
db, err := s.requireDB()
|
||||
if err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
ctx, cancel := s.beginOperation(opts.TimeoutSecs)
|
||||
defer s.endOperation(cancel)
|
||||
execResult, err := db.ExecContext(ctx, sqlText)
|
||||
if err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
affected, _ := execResult.RowsAffected()
|
||||
return queryResult{Columns: []string{}, ColumnTypes: []string{}, Rows: [][]any{}, AffectedRows: affected, ExecutionTimeMS: time.Since(start).Milliseconds()}, nil
|
||||
}
|
||||
|
||||
func (s *server) queryRows(sqlText string, timeoutSecs int) (*sql.Rows, context.CancelFunc, error) {
|
||||
db, err := s.requireDB()
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
ctx, cancel := s.beginOperation(timeoutSecs)
|
||||
rows, err := db.QueryContext(ctx, sqlText)
|
||||
if err != nil {
|
||||
s.endOperation(cancel)
|
||||
return nil, nil, err
|
||||
}
|
||||
return rows, cancel, nil
|
||||
}
|
||||
|
||||
func (s *server) executeQueryPage(opts queryOptions, pageSize int) (queryPageResult, error) {
|
||||
start := time.Now()
|
||||
if err := s.setSchema(opts.Schema); err != nil {
|
||||
return queryPageResult{}, err
|
||||
}
|
||||
sqlText := trimStatementSQL(opts.SQL)
|
||||
if !isQuerySQL(sqlText) {
|
||||
result, err := s.executeQuery(opts)
|
||||
return queryPageResult{Columns: result.Columns, ColumnTypes: result.ColumnTypes, Rows: result.Rows, AffectedRows: result.AffectedRows, ExecutionTimeMS: result.ExecutionTimeMS, Truncated: result.Truncated}, err
|
||||
}
|
||||
rows, cancel, err := s.queryRows(sqlText, opts.TimeoutSecs)
|
||||
if err != nil {
|
||||
return queryPageResult{}, err
|
||||
}
|
||||
columns, err := rows.Columns()
|
||||
if err != nil {
|
||||
_ = rows.Close()
|
||||
s.endOperation(cancel)
|
||||
return queryPageResult{}, err
|
||||
}
|
||||
maxRows := opts.MaxRows
|
||||
if maxRows <= 0 {
|
||||
maxRows = defaultMaxRows
|
||||
}
|
||||
session := &querySession{rows: rows, columns: columns, columnTypes: columnTypeNames(rows), remaining: maxRows, cancel: cancel}
|
||||
result, err := readQuerySessionPage(session, pageSize)
|
||||
result.ExecutionTimeMS = time.Since(start).Milliseconds()
|
||||
if err != nil {
|
||||
_ = rows.Close()
|
||||
s.endOperation(cancel)
|
||||
return queryPageResult{}, err
|
||||
}
|
||||
if result.HasMore {
|
||||
s.nextSessionID++
|
||||
id := fmt.Sprintf("kingbase-%d", s.nextSessionID)
|
||||
s.sessions[id] = session
|
||||
result.SessionID = &id
|
||||
} else {
|
||||
_ = rows.Close()
|
||||
s.endOperation(cancel)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *server) fetchQueryPage(id string, pageSize int) (queryPageResult, error) {
|
||||
session := s.sessions[id]
|
||||
if session == nil {
|
||||
return queryPageResult{Columns: []string{}, ColumnTypes: []string{}, Rows: [][]any{}}, nil
|
||||
}
|
||||
result, err := readQuerySessionPage(session, pageSize)
|
||||
if err != nil {
|
||||
s.closeQuerySession(id)
|
||||
return queryPageResult{}, err
|
||||
}
|
||||
if result.HasMore {
|
||||
result.SessionID = &id
|
||||
} else {
|
||||
s.closeQuerySession(id)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *server) closeQuerySession(id string) bool {
|
||||
session := s.sessions[id]
|
||||
if session == nil {
|
||||
return false
|
||||
}
|
||||
_ = session.rows.Close()
|
||||
if session.cancel != nil {
|
||||
s.endOperation(session.cancel)
|
||||
}
|
||||
delete(s.sessions, id)
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *server) closeAllQuerySessions() {
|
||||
for id := range s.sessions {
|
||||
s.closeQuerySession(id)
|
||||
}
|
||||
}
|
||||
|
||||
func readQuerySessionPage(session *querySession, pageSize int) (queryPageResult, error) {
|
||||
if pageSize <= 0 {
|
||||
pageSize = 100
|
||||
}
|
||||
capacity := min(pageSize, session.remaining)
|
||||
result := queryPageResult{Columns: session.columns, ColumnTypes: session.columnTypes, Rows: make([][]any, 0, capacity)}
|
||||
for len(result.Rows) < pageSize && session.remaining > 0 {
|
||||
if session.pending != nil {
|
||||
result.Rows = append(result.Rows, session.pending)
|
||||
session.pending = nil
|
||||
session.remaining--
|
||||
continue
|
||||
}
|
||||
if !session.rows.Next() {
|
||||
return result, session.rows.Err()
|
||||
}
|
||||
row, err := scanRow(session.rows, len(session.columns))
|
||||
if err != nil {
|
||||
return queryPageResult{}, err
|
||||
}
|
||||
result.Rows = append(result.Rows, row)
|
||||
session.remaining--
|
||||
}
|
||||
if session.remaining <= 0 {
|
||||
result.Truncated = true
|
||||
return result, nil
|
||||
}
|
||||
if session.rows.Next() {
|
||||
row, err := scanRow(session.rows, len(session.columns))
|
||||
if err != nil {
|
||||
return queryPageResult{}, err
|
||||
}
|
||||
session.pending = row
|
||||
result.HasMore = true
|
||||
}
|
||||
return result, session.rows.Err()
|
||||
}
|
||||
|
||||
func readRows(rows *sql.Rows, maxRows int) (queryResult, error) {
|
||||
columns, err := rows.Columns()
|
||||
if err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
result := queryResult{Columns: columns, ColumnTypes: columnTypeNames(rows), Rows: make([][]any, 0, min(maxRows, 1024))}
|
||||
for rows.Next() {
|
||||
if len(result.Rows) >= maxRows {
|
||||
result.Truncated = true
|
||||
break
|
||||
}
|
||||
row, err := scanRow(rows, len(columns))
|
||||
if err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
result.Rows = append(result.Rows, row)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func scanRow(rows *sql.Rows, count int) ([]any, error) {
|
||||
storage := make([]any, count*2)
|
||||
values := storage[:count]
|
||||
dest := storage[count:]
|
||||
for i := range values {
|
||||
dest[i] = &values[i]
|
||||
}
|
||||
if err := rows.Scan(dest...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i, value := range values {
|
||||
values[i] = normalizeValue(value)
|
||||
}
|
||||
return values, nil
|
||||
}
|
||||
|
||||
func columnTypeNames(rows *sql.Rows) []string {
|
||||
types, err := rows.ColumnTypes()
|
||||
if err != nil {
|
||||
return []string{}
|
||||
}
|
||||
result := make([]string, len(types))
|
||||
for i, columnType := range types {
|
||||
result[i] = columnType.DatabaseTypeName()
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (s *server) executeTransaction(params map[string]json.RawMessage) (queryResult, error) {
|
||||
db, err := s.requireDB()
|
||||
if err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
statements := stringSliceParam(params, "statements")
|
||||
if err := s.setSchema(stringParam(params, "schema")); err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
start := time.Now()
|
||||
tx, err := db.Begin()
|
||||
if err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
var affected int64
|
||||
for _, statement := range statements {
|
||||
result, execErr := tx.Exec(trimStatementSQL(statement))
|
||||
if execErr != nil {
|
||||
_ = tx.Rollback()
|
||||
return queryResult{}, execErr
|
||||
}
|
||||
rows, _ := result.RowsAffected()
|
||||
affected += rows
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
return queryResult{Columns: []string{}, ColumnTypes: []string{}, Rows: [][]any{}, AffectedRows: affected, ExecutionTimeMS: time.Since(start).Milliseconds()}, nil
|
||||
}
|
||||
|
||||
func (s *server) executeBatch(params map[string]json.RawMessage) (queryResult, error) {
|
||||
start := time.Now()
|
||||
var affected int64
|
||||
for _, statement := range stringSliceParam(params, "statements") {
|
||||
result, err := s.executeQuery(queryOptions{SQL: statement, Schema: stringParam(params, "schema")})
|
||||
if err != nil {
|
||||
return queryResult{}, err
|
||||
}
|
||||
affected += result.AffectedRows
|
||||
}
|
||||
return queryResult{Columns: []string{}, ColumnTypes: []string{}, Rows: [][]any{}, AffectedRows: affected, ExecutionTimeMS: time.Since(start).Milliseconds()}, nil
|
||||
}
|
||||
|
||||
func (s *server) setSchema(schema string) error {
|
||||
schema = strings.TrimSpace(schema)
|
||||
if schema == "" && !s.schemaSet {
|
||||
return nil
|
||||
}
|
||||
if schema != "" && s.schemaSet && schema == s.currentSchema {
|
||||
return nil
|
||||
}
|
||||
db, err := s.requireDB()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
statement := "RESET search_path"
|
||||
if schema != "" {
|
||||
// Kingbase implicitly prioritizes its system catalog when it is not
|
||||
// listed explicitly, matching the JDBC agent and DBeaver behavior.
|
||||
statement = "SET search_path TO " + quoteIdentifier(schema)
|
||||
}
|
||||
if _, err = db.Exec(statement); err != nil {
|
||||
return err
|
||||
}
|
||||
s.currentSchema = schema
|
||||
s.schemaSet = schema != ""
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildDSN(cp connectParams) string {
|
||||
if value := strings.TrimSpace(cp.ConnectionString); value != "" && !isKingbaseJDBCURL(value) {
|
||||
return value
|
||||
}
|
||||
port := cp.Port
|
||||
if port <= 0 {
|
||||
port = 54321
|
||||
}
|
||||
sslMode := "disable"
|
||||
if cp.SSL {
|
||||
sslMode = "verify-full"
|
||||
}
|
||||
parts := []string{
|
||||
"host=" + quoteDSNValue(cp.Host),
|
||||
fmt.Sprintf("port=%d", port),
|
||||
"user=" + quoteDSNValue(cp.Username),
|
||||
"password=" + quoteDSNValue(cp.Password),
|
||||
"dbname=" + quoteDSNValue(cp.Database),
|
||||
"sslmode=" + sslMode,
|
||||
"connect_timeout=15",
|
||||
}
|
||||
if cp.CACertPath != "" {
|
||||
parts = append(parts, "sslrootcert="+quoteDSNValue(cp.CACertPath))
|
||||
}
|
||||
if cp.ClientCertPath != "" {
|
||||
parts = append(parts, "sslcert="+quoteDSNValue(cp.ClientCertPath))
|
||||
}
|
||||
if cp.ClientKeyPath != "" {
|
||||
parts = append(parts, "sslkey="+quoteDSNValue(cp.ClientKeyPath))
|
||||
}
|
||||
for _, pair := range strings.FieldsFunc(cp.URLParams, func(r rune) bool { return r == '&' || r == ';' }) {
|
||||
key, value, ok := strings.Cut(pair, "=")
|
||||
if ok && isSafeParamKey(key) {
|
||||
parts = append(parts, strings.TrimSpace(key)+"="+quoteDSNValue(strings.TrimSpace(value)))
|
||||
}
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func isKingbaseJDBCURL(value string) bool {
|
||||
return strings.HasPrefix(strings.ToLower(strings.TrimSpace(value)), "jdbc:kingbase8://")
|
||||
}
|
||||
|
||||
func quoteDSNValue(value string) string {
|
||||
return "'" + strings.ReplaceAll(strings.ReplaceAll(value, `\`, `\\`), "'", `\'`) + "'"
|
||||
}
|
||||
|
||||
func isSafeParamKey(value string) bool {
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" {
|
||||
return false
|
||||
}
|
||||
for _, char := range value {
|
||||
if !(char == '_' || char >= 'a' && char <= 'z' || char >= 'A' && char <= 'Z' || char >= '0' && char <= '9') {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func normalizeValue(value any) any {
|
||||
switch typed := value.(type) {
|
||||
case nil:
|
||||
return nil
|
||||
case []byte:
|
||||
if isTextBytes(typed) {
|
||||
return string(typed)
|
||||
}
|
||||
return map[string]string{"$binary": base64.StdEncoding.EncodeToString(typed)}
|
||||
case time.Time:
|
||||
return typed.Format(time.RFC3339Nano)
|
||||
case int8:
|
||||
return int64(typed)
|
||||
case int16:
|
||||
return int64(typed)
|
||||
case int32:
|
||||
return int64(typed)
|
||||
case float32:
|
||||
return float64(typed)
|
||||
default:
|
||||
return typed
|
||||
}
|
||||
}
|
||||
|
||||
func isTextBytes(value []byte) bool {
|
||||
for _, char := range value {
|
||||
if char == 0 || char < 0x09 || char > 0x0d && char < 0x20 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func decodeParams(params map[string]json.RawMessage, target any) error {
|
||||
data, err := json.Marshal(params)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return json.Unmarshal(data, target)
|
||||
}
|
||||
|
||||
func stringParam(params map[string]json.RawMessage, key string) string {
|
||||
var value string
|
||||
_ = json.Unmarshal(params[key], &value)
|
||||
return value
|
||||
}
|
||||
|
||||
func intParam(params map[string]json.RawMessage, key string) int {
|
||||
var value int
|
||||
_ = json.Unmarshal(params[key], &value)
|
||||
return value
|
||||
}
|
||||
|
||||
func stringSliceParam(params map[string]json.RawMessage, key string) []string {
|
||||
var values []string
|
||||
if json.Unmarshal(params[key], &values) == nil {
|
||||
return values
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func metadataListConstraintsFromParams(params map[string]json.RawMessage) metadataListConstraints {
|
||||
return metadataListConstraints{
|
||||
Filter: stringParam(params, "filter"),
|
||||
Limit: intParam(params, "limit"),
|
||||
Offset: intParam(params, "offset"),
|
||||
ObjectTypes: stringSliceParam(params, "object_types"),
|
||||
}
|
||||
}
|
||||
|
||||
func errorResponse(id json.RawMessage, err error) response {
|
||||
return response{JSONRPC: "2.0", ID: id, Error: &rpcError{Code: -1, Message: err.Error()}}
|
||||
}
|
||||
|
||||
func trimStatementSQL(sqlText string) string {
|
||||
return strings.TrimRight(strings.TrimSpace(sqlText), "; \t\r\n")
|
||||
}
|
||||
|
||||
func isQuerySQL(sqlText string) bool {
|
||||
lower := strings.ToLower(strings.TrimSpace(sqlText))
|
||||
return strings.HasPrefix(lower, "select") || strings.HasPrefix(lower, "with") || strings.HasPrefix(lower, "show") || strings.HasPrefix(lower, "explain")
|
||||
}
|
||||
|
||||
func quoteIdentifier(value string) string {
|
||||
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
|
||||
}
|
||||
|
||||
func quoteLiteral(value string) string {
|
||||
return "'" + strings.ReplaceAll(value, "'", "''") + "'"
|
||||
}
|
||||
|
||||
func stringPtr(value string) *string {
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
return &value
|
||||
}
|
||||
|
|
@ -0,0 +1,359 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"database/sql/driver"
|
||||
"errors"
|
||||
"io"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gitea.com/kingbase/gokb"
|
||||
)
|
||||
|
||||
var registerTestDriver sync.Once
|
||||
var testDriverState atomic.Pointer[fakeDriverState]
|
||||
var registerExpressionFallbackDriver sync.Once
|
||||
var expressionFallbackState atomic.Pointer[fallbackDriverState]
|
||||
|
||||
type fakeDriverState struct {
|
||||
queryArgs int
|
||||
queryCtx context.Context
|
||||
rowCount int
|
||||
}
|
||||
|
||||
type fakeDriver struct{}
|
||||
|
||||
type fakeConn struct{}
|
||||
|
||||
type fakeRows struct {
|
||||
current int
|
||||
count int
|
||||
}
|
||||
|
||||
type fallbackDriverState struct {
|
||||
mu sync.Mutex
|
||||
queries []string
|
||||
}
|
||||
|
||||
type fallbackDriver struct{}
|
||||
|
||||
type fallbackConn struct {
|
||||
state *fallbackDriverState
|
||||
}
|
||||
|
||||
type valueRows struct {
|
||||
columns []string
|
||||
rows [][]driver.Value
|
||||
index int
|
||||
}
|
||||
|
||||
func (fakeDriver) Open(string) (driver.Conn, error) { return fakeConn{}, nil }
|
||||
|
||||
func (fakeConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip }
|
||||
|
||||
func (fakeConn) Close() error { return nil }
|
||||
|
||||
func (fakeConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip }
|
||||
|
||||
func (fakeConn) QueryContext(ctx context.Context, _ string, args []driver.NamedValue) (driver.Rows, error) {
|
||||
state := testDriverState.Load()
|
||||
state.queryArgs = len(args)
|
||||
state.queryCtx = ctx
|
||||
return &fakeRows{count: state.rowCount}, nil
|
||||
}
|
||||
|
||||
func (fakeConn) ExecContext(context.Context, string, []driver.NamedValue) (driver.Result, error) {
|
||||
return driver.RowsAffected(1), nil
|
||||
}
|
||||
|
||||
func (fakeRows) Columns() []string { return []string{"value"} }
|
||||
|
||||
func (fakeRows) Close() error { return nil }
|
||||
|
||||
func (rows *fakeRows) Next(values []driver.Value) error {
|
||||
if rows.current >= rows.count {
|
||||
return io.EOF
|
||||
}
|
||||
rows.current++
|
||||
values[0] = int64(rows.current)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (fallbackDriver) Open(string) (driver.Conn, error) {
|
||||
return &fallbackConn{state: expressionFallbackState.Load()}, nil
|
||||
}
|
||||
|
||||
func (*fallbackConn) Prepare(string) (driver.Stmt, error) { return nil, driver.ErrSkip }
|
||||
|
||||
func (*fallbackConn) Close() error { return nil }
|
||||
|
||||
func (*fallbackConn) Begin() (driver.Tx, error) { return nil, driver.ErrSkip }
|
||||
|
||||
func (connection *fallbackConn) QueryContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Rows, error) {
|
||||
connection.state.mu.Lock()
|
||||
connection.state.queries = append(connection.state.queries, query)
|
||||
connection.state.mu.Unlock()
|
||||
if strings.Contains(query, "information_schema.table_constraints") {
|
||||
return &valueRows{columns: []string{"column_name"}}, nil
|
||||
}
|
||||
if strings.Contains(query, "sys_get_expr(") {
|
||||
return nil, &gokb.Error{Code: gokb.ErrorCode("42883"), Message: "function sys_get_expr(pg_node_tree, oid) does not exist"}
|
||||
}
|
||||
if strings.Contains(query, "pg_get_expr(") {
|
||||
return &valueRows{
|
||||
columns: []string{"column_name", "data_type", "is_nullable", "column_default", "column_comment", "numeric_precision", "numeric_scale", "character_maximum_length"},
|
||||
rows: [][]driver.Value{{"id", "integer", false, "nextval('orders_id_seq'::regclass)", nil, int64(32), int64(0), nil}},
|
||||
}, nil
|
||||
}
|
||||
return nil, errors.New("unexpected query: " + query)
|
||||
}
|
||||
|
||||
func (rows *valueRows) Columns() []string { return rows.columns }
|
||||
|
||||
func (*valueRows) Close() error { return nil }
|
||||
|
||||
func (rows *valueRows) Next(values []driver.Value) error {
|
||||
if rows.index >= len(rows.rows) {
|
||||
return io.EOF
|
||||
}
|
||||
copy(values, rows.rows[rows.index])
|
||||
rows.index++
|
||||
return nil
|
||||
}
|
||||
|
||||
func openFakeDB(t *testing.T, rowCount int) (*sql.DB, *fakeDriverState) {
|
||||
t.Helper()
|
||||
registerTestDriver.Do(func() { sql.Register("kingbase-agent-test", fakeDriver{}) })
|
||||
state := &fakeDriverState{rowCount: rowCount}
|
||||
testDriverState.Store(state)
|
||||
db, err := sql.Open("kingbase-agent-test", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db.SetMaxOpenConns(1)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
return db, state
|
||||
}
|
||||
|
||||
func TestHandshakeAdvertisesMultiSession(t *testing.T) {
|
||||
runtime := &runtimeServer{sessions: map[string]*agentSession{}}
|
||||
result, shutdown, err := runtime.dispatch("handshake", nil)
|
||||
if err != nil || shutdown {
|
||||
t.Fatalf("handshake failed: shutdown=%v err=%v", shutdown, err)
|
||||
}
|
||||
values := result.(map[string]any)
|
||||
if values["protocolVersion"] != protocolVersion {
|
||||
t.Fatalf("unexpected protocol version: %#v", values["protocolVersion"])
|
||||
}
|
||||
capabilities := values["capabilities"].([]string)
|
||||
if !containsString(capabilities, "multi_session") || !containsString(capabilities, "paged_query") {
|
||||
t.Fatalf("missing capabilities: %v", capabilities)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildDSNQuotesCredentialsAndFiltersKeys(t *testing.T) {
|
||||
dsn := buildDSN(connectParams{
|
||||
Host: "db host",
|
||||
Port: 54321,
|
||||
Database: "test'db",
|
||||
Username: "system",
|
||||
Password: `p'ass\\word`,
|
||||
URLParams: "application_name=dbx&bad-key=ignored",
|
||||
})
|
||||
for _, expected := range []string{
|
||||
`host='db host'`, `dbname='test\'db'`, `password='p\'ass\\\\word'`, `application_name='dbx'`,
|
||||
} {
|
||||
if !strings.Contains(dsn, expected) {
|
||||
t.Fatalf("DSN missing %q: %s", expected, dsn)
|
||||
}
|
||||
}
|
||||
if strings.Contains(dsn, "bad-key") {
|
||||
t.Fatalf("unsafe parameter key was accepted: %s", dsn)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildDSNConvertsDBXJDBCURL(t *testing.T) {
|
||||
dsn := buildDSN(connectParams{
|
||||
Host: "127.0.0.1",
|
||||
Port: 54321,
|
||||
Database: "test",
|
||||
Username: "system",
|
||||
Password: "secret",
|
||||
URLParams: "application_name=dbx",
|
||||
ConnectionString: "jdbc:kingbase8://127.0.0.1:54321/test?application_name=dbx",
|
||||
})
|
||||
if strings.HasPrefix(dsn, "jdbc:") || !strings.Contains(dsn, "host='127.0.0.1'") || !strings.Contains(dsn, "dbname='test'") {
|
||||
t.Fatalf("JDBC URL was not converted to a gokb DSN: %s", dsn)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMetadataNormalizationHelpers(t *testing.T) {
|
||||
if normalizeTableType("BASE TABLE") != "TABLE" {
|
||||
t.Fatal("BASE TABLE was not normalized")
|
||||
}
|
||||
if decodeTriggerTiming(1<<6) != "INSTEAD OF" || decodeTriggerTiming(1<<1) != "BEFORE" || decodeTriggerTiming(0) != "AFTER" {
|
||||
t.Fatal("trigger timing decoding is incorrect")
|
||||
}
|
||||
length := boundedVarcharLength("character varying ( 128 )")
|
||||
if length == nil || *length != 128 {
|
||||
t.Fatalf("bounded varchar length not parsed: %v", length)
|
||||
}
|
||||
if boundedVarcharLength("text") != nil {
|
||||
t.Fatal("unbounded type returned a length")
|
||||
}
|
||||
}
|
||||
|
||||
func TestColumnsFallbackToPgGetExprAndCacheChoice(t *testing.T) {
|
||||
registerExpressionFallbackDriver.Do(func() { sql.Register("kingbase-expression-fallback-test", fallbackDriver{}) })
|
||||
state := &fallbackDriverState{}
|
||||
expressionFallbackState.Store(state)
|
||||
db, err := sql.Open("kingbase-expression-fallback-test", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db.SetMaxOpenConns(1)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
server := newServer()
|
||||
server.db = db
|
||||
|
||||
for call := 0; call < 2; call++ {
|
||||
columns, err := server.getColumns("public", "orders")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(columns) != 1 || columns[0].ColumnDefault == nil || *columns[0].ColumnDefault != "nextval('orders_id_seq'::regclass)" {
|
||||
t.Fatalf("unexpected columns: %#v", columns)
|
||||
}
|
||||
}
|
||||
state.mu.Lock()
|
||||
defer state.mu.Unlock()
|
||||
var sysCalls, pgCalls int
|
||||
for _, query := range state.queries {
|
||||
if strings.Contains(query, "sys_get_expr(") {
|
||||
sysCalls++
|
||||
}
|
||||
if strings.Contains(query, "pg_get_expr(") {
|
||||
pgCalls++
|
||||
}
|
||||
}
|
||||
if sysCalls != 1 || pgCalls != 2 {
|
||||
t.Fatalf("fallback choice was not cached: sys=%d pg=%d queries=%v", sysCalls, pgCalls, state.queries)
|
||||
}
|
||||
}
|
||||
|
||||
func TestQuoteLiteralEscapesMetadataValues(t *testing.T) {
|
||||
if got := quoteLiteral("a'b"); got != "'a''b'" {
|
||||
t.Fatalf("unexpected literal: %s", got)
|
||||
}
|
||||
constraints := metadataListConstraints{Filter: "CHILD", ObjectTypes: []string{"table"}}
|
||||
if !constraintsMatch(constraints, "dbx_child", "TABLE") || constraintsMatch(constraints, "dbx_parent", "TABLE") {
|
||||
t.Fatal("metadata constraints were not applied")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompletionNameMatching(t *testing.T) {
|
||||
request := completionAssistantRequest{Mask: "DBX_", MatchMode: "prefix"}
|
||||
if !completionNameMatches("dbx_child", request) || completionNameMatches("other_dbx_child", request) {
|
||||
t.Fatal("case-insensitive prefix matching failed")
|
||||
}
|
||||
request.MatchMode = "contains"
|
||||
if !completionNameMatches("other_dbx_child", request) {
|
||||
t.Fatal("contains matching failed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecuteQueryUsesSimpleProtocolAndReleasesContext(t *testing.T) {
|
||||
db, state := openFakeDB(t, 1)
|
||||
server := newServer()
|
||||
server.db = db
|
||||
result, err := server.executeQuery(queryOptions{SQL: "SELECT 1", MaxRows: 10})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(result.Rows) != 1 || state.queryArgs != 0 {
|
||||
t.Fatalf("unexpected result or bound arguments: rows=%v args=%d", result.Rows, state.queryArgs)
|
||||
}
|
||||
assertContextCanceled(t, state.queryCtx)
|
||||
}
|
||||
|
||||
func TestPagedQueryKeepsContextUntilSessionCloses(t *testing.T) {
|
||||
db, state := openFakeDB(t, 3)
|
||||
server := newServer()
|
||||
server.db = db
|
||||
result, err := server.executeQueryPage(queryOptions{SQL: "SELECT value", MaxRows: 10}, 1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.HasMore || result.SessionID == nil {
|
||||
t.Fatalf("expected an open query session: %#v", result)
|
||||
}
|
||||
select {
|
||||
case <-state.queryCtx.Done():
|
||||
t.Fatal("paged query context was canceled before session close")
|
||||
default:
|
||||
}
|
||||
if !server.closeQuerySession(*result.SessionID) {
|
||||
t.Fatal("query session was not closed")
|
||||
}
|
||||
assertContextCanceled(t, state.queryCtx)
|
||||
}
|
||||
|
||||
func TestRuntimeCloseSessionWaitsForActiveRequestAndClosesTarget(t *testing.T) {
|
||||
db, _ := openFakeDB(t, 0)
|
||||
target := &agentSession{server: newServer()}
|
||||
target.server.db = db
|
||||
other := &agentSession{server: newServer()}
|
||||
runtime := &runtimeServer{sessions: map[string]*agentSession{"target": target, "other": other}}
|
||||
|
||||
target.mu.Lock()
|
||||
closed := make(chan error, 1)
|
||||
go func() { closed <- runtime.closeSession("target") }()
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
select {
|
||||
case err := <-closed:
|
||||
t.Fatalf("close_session returned before the active request completed: %v", err)
|
||||
default:
|
||||
}
|
||||
if _, err := runtime.session("target"); err == nil {
|
||||
t.Fatal("draining session remained available for new requests")
|
||||
}
|
||||
if _, err := runtime.session("other"); err != nil {
|
||||
t.Fatalf("unrelated session was removed: %v", err)
|
||||
}
|
||||
|
||||
target.mu.Unlock()
|
||||
select {
|
||||
case err := <-closed:
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("close_session did not finish after the active request released the session")
|
||||
}
|
||||
if target.server.db != nil {
|
||||
t.Fatal("target database connection was not closed")
|
||||
}
|
||||
}
|
||||
|
||||
func assertContextCanceled(t *testing.T, ctx context.Context) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("query context was not canceled")
|
||||
}
|
||||
}
|
||||
|
||||
func containsString(values []string, target string) bool {
|
||||
for _, value := range values {
|
||||
if value == target {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
|
@ -47,6 +47,8 @@ for platform in "${PLATFORMS[@]}"; do
|
|||
# Copy all driver JARs (platform-independent)
|
||||
for jar_file in "$RELEASE_DIR"/dbx-agent-*.jar; do
|
||||
[ -f "$jar_file" ] || continue
|
||||
# Kingbase is distributed only as a native agent; keep legacy JDBC builds out of offline bundles.
|
||||
[ "$(basename "$jar_file")" = "dbx-agent-kingbase.jar" ] && continue
|
||||
cp "$jar_file" "$WORK/drivers/"
|
||||
done
|
||||
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ const i18n = {
|
|||
drivers: "Database Drivers",
|
||||
driversDesc: "JDBC driver JAR files for each supported database type.",
|
||||
nativeAgents: "Native Agents",
|
||||
nativeAgentsDesc: "Go-based native agents for Oracle and XuguDB. Download the executable that matches the offline machine.",
|
||||
nativeAgentsDesc: "Go-based native agents for Oracle, KingBase, and XuguDB. Download the executable that matches the offline machine.",
|
||||
jre: "Java Runtime (JRE)",
|
||||
jreDesc: "JRE packages used by agent-based database drivers. Required for Oracle, SQL Server, and other agent-managed connections.",
|
||||
loading: "Loading driver catalog...",
|
||||
|
|
@ -58,7 +58,7 @@ const i18n = {
|
|||
drivers: "数据库驱动",
|
||||
driversDesc: "每种支持的数据库类型对应的 JDBC 驱动 JAR 文件。",
|
||||
nativeAgents: "原生 Agent",
|
||||
nativeAgentsDesc: "Oracle 和虚谷使用 Go 原生 Agent,请下载与内网机器平台匹配的可执行文件。",
|
||||
nativeAgentsDesc: "Oracle、人大金仓和虚谷使用 Go 原生 Agent,请下载与内网机器平台匹配的可执行文件。",
|
||||
jre: "Java 运行时 (JRE)",
|
||||
jreDesc: "Agent 驱动所需的 JRE 环境,Oracle、SQL Server 等数据库通过 Agent 连接时需要。",
|
||||
loading: "正在加载驱动列表...",
|
||||
|
|
|
|||
|
|
@ -339,7 +339,7 @@ agents/drivers/<驱动模块名>/build/libs/
|
|||
|
||||
#### 原生 Agent
|
||||
|
||||
`oracle`、`xugu` 等原生 Agent 使用 `agent` 可执行文件,而不是 `agent.jar`。在对应模块目录运行 Go 测试和构建,并按模块 README 替换本地运行时可执行文件:
|
||||
`oracle`、`kingbase`、`xugu` 等原生 Agent 使用 `agent` 可执行文件,而不是 `agent.jar`。在对应模块目录运行 Go 测试和构建,并按模块 README 替换本地运行时可执行文件:
|
||||
|
||||
```bash
|
||||
cd agents/drivers/<原生模块>
|
||||
|
|
|
|||
|
|
@ -339,7 +339,7 @@ Normal bug-fix pull requests do not manually bump a version to make the change t
|
|||
|
||||
#### Native Agents
|
||||
|
||||
Native agents such as `oracle` and `xugu` use an `agent` executable instead of `agent.jar`. Run the Go tests and build from the module, then follow its README to replace the local runtime executable:
|
||||
Native agents such as `oracle`, `kingbase`, and `xugu` use an `agent` executable instead of `agent.jar`. Run the Go tests and build from the module, then follow its README to replace the local runtime executable:
|
||||
|
||||
```bash
|
||||
cd agents/drivers/<native-module>
|
||||
|
|
|
|||
|
|
@ -1,18 +1,19 @@
|
|||
---
|
||||
title: 驱动管理
|
||||
description: 管理内置和 Agent 驱动的 JDBC 驱动,配置 JRE 版本,处理驱动更新。
|
||||
description: 管理内置驱动、原生 Agent 和 JDBC Agent,配置 JRE 版本并处理驱动更新。
|
||||
---
|
||||
|
||||
DBX 使用混合驱动架构:常用数据库使用内置原生驱动,需要供应商特定 JDBC 驱动的数据库使用 JDBC Agent 系统。
|
||||
DBX 使用混合驱动架构:内置 Rust 驱动、独立原生 Agent,以及 Java/JDBC Agent。
|
||||
|
||||
## 驱动架构
|
||||
|
||||
| 驱动类型 | 工作方式 | 适用场景 |
|
||||
| ------------ | ---------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
|
||||
| 原生(内置) | Rust 驱动编译进 DBX | MySQL、PostgreSQL、SQLite、SQL Server、Oracle、Redis、MongoDB、DuckDB、ClickHouse 等 |
|
||||
| JDBC Agent | DBX 管理的 Java 子进程 | 提供 JDBC 驱动的数据库:GaussDB、openGauss、DM、KingBase、HighGo、Vastbase、Trino、Hive、DB2、Informix、Neo4j、TDengine、虚谷 XuguDB、YashanDB、GoldenDB、Kylin、SunDB 等 |
|
||||
| 驱动类型 | 工作方式 | 适用场景 |
|
||||
| ------------ | --------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------- |
|
||||
| 原生(内置) | Rust 驱动编译进 DBX | MySQL、PostgreSQL、SQLite、SQL Server、Redis、MongoDB、DuckDB、ClickHouse 等 |
|
||||
| 原生 Agent | DBX 管理的 Go/Rust 独立进程 | Oracle、KingBase、虚谷 XuguDB,以及存在成熟原生驱动的数据库 |
|
||||
| JDBC Agent | DBX 管理的 Java 子进程 | 需要 JDBC 的数据库:GaussDB、openGauss、DM、HighGo、Vastbase、Trino、Hive、DB2、Informix、Neo4j、TDengine、YashanDB、GoldenDB、Kylin、SunDB 等 |
|
||||
|
||||
<Callout type="info">原生驱动安装后立即可用。Agent 驱动需要一次性下载 JDBC 驱动和 JRE 设置,首次创建连接时 DBX 会自动处理。</Callout>
|
||||
<Callout type="info">内置驱动可直接使用;原生 Agent 只需对应平台的可执行文件,JDBC Agent 还需要 JRE。首次创建连接时 DBX 会自动安装匹配的组件。</Callout>
|
||||
|
||||
## 驱动商店
|
||||
|
||||
|
|
@ -29,7 +30,7 @@ DBX 使用混合驱动架构:常用数据库使用内置原生驱动,需要
|
|||
<Steps>
|
||||
<Step>### 打开驱动商店 导航到**设置 → 驱动**或点击连接对话框中的驱动安装提示。</Step>
|
||||
<Step>### 选择驱动 找到所需的数据库驱动。每个条目显示支持的数据库和驱动版本。</Step>
|
||||
<Step>### 点击安装 DBX 会下载 JDBC 驱动 JAR 及所需的依赖项。下载过程中显示进度。</Step>
|
||||
<Step>### 点击安装 DBX 会下载原生 Agent 可执行文件或 JDBC Agent JAR,以及所需运行时。下载过程中显示进度。</Step>
|
||||
<Step>### 创建连接 返回连接对话框。驱动此时即可使用。</Step>
|
||||
</Steps>
|
||||
|
||||
|
|
|
|||
|
|
@ -1,18 +1,19 @@
|
|||
---
|
||||
title: Driver Management
|
||||
description: Manage built-in and agent-based JDBC drivers, configure JRE versions, and handle driver updates.
|
||||
description: Manage built-in, native Agent, and JDBC Agent drivers, configure JRE versions, and handle driver updates.
|
||||
---
|
||||
|
||||
DBX uses a hybrid driver architecture: built-in native drivers for common databases and a JDBC agent system for databases that require vendor-specific JDBC drivers.
|
||||
DBX uses a hybrid driver architecture: built-in Rust drivers, standalone native Agents, and Java/JDBC Agents.
|
||||
|
||||
## Driver Architecture
|
||||
|
||||
| Driver Type | How It Works | Best For |
|
||||
| ----------------- | ------------------------------ | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
|
||||
| Native (built-in) | Rust drivers compiled into DBX | MySQL, PostgreSQL, SQLite, SQL Server, Oracle, Redis, MongoDB, DuckDB, ClickHouse, and more |
|
||||
| JDBC Agent | Java subprocess managed by DBX | Databases that provide JDBC drivers: GaussDB, openGauss, DM, KingBase, HighGo, Vastbase, Trino, Hive, DB2, Informix, Neo4j, TDengine, XuguDB, YashanDB, GoldenDB, Kylin, SunDB, and more |
|
||||
| Driver Type | How It Works | Best For |
|
||||
| ----------------- | ------------------------------------ | --------------------------------------------------------------------------------------------------------------------------------------------------------------- |
|
||||
| Native (built-in) | Rust drivers compiled into DBX | MySQL, PostgreSQL, SQLite, SQL Server, Redis, MongoDB, DuckDB, ClickHouse, and more |
|
||||
| Native Agent | Standalone Go/Rust process managed by DBX | Oracle, KingBase, XuguDB, and databases with a mature native driver |
|
||||
| JDBC Agent | Java subprocess managed by DBX | Databases that require JDBC: GaussDB, openGauss, DM, HighGo, Vastbase, Trino, Hive, DB2, Informix, Neo4j, TDengine, YashanDB, GoldenDB, Kylin, SunDB, and more |
|
||||
|
||||
<Callout type="info">Native drivers work immediately after installation. Agent drivers require a one-time JDBC driver download and JRE setup, which DBX handles automatically when you first create a connection.</Callout>
|
||||
<Callout type="info">Built-in drivers work immediately. Native Agents require only a platform-specific executable, while JDBC Agents also require a JRE. DBX handles the matching installation when you first create a connection.</Callout>
|
||||
|
||||
## Driver Store
|
||||
|
||||
|
|
@ -29,7 +30,7 @@ Open the Driver Store from **Settings → Drivers** or click the driver hint tha
|
|||
<Steps>
|
||||
<Step>### Open Driver Store Navigate to **Settings → Drivers** or click the driver install hint in the connection dialog.</Step>
|
||||
<Step>### Choose a Driver Find the database driver you need. Each entry shows the supported database and driver version.</Step>
|
||||
<Step>### Click Install DBX downloads the JDBC driver JAR and any required dependencies. Progress is shown during download.</Step>
|
||||
<Step>### Click Install DBX downloads the native Agent executable or JDBC Agent JAR and any required runtime. Progress is shown during download.</Step>
|
||||
<Step>### Create a Connection Return to the connection dialog. The driver is now ready for use.</Step>
|
||||
</Steps>
|
||||
|
||||
|
|
|
|||
|
|
@ -105,7 +105,7 @@ PostgreSQL、MySQL、SQLite、兼容 SQL 数据库、独立 Redis 和 MongoDB
|
|||
|
||||
### Agent/JDBC 数据库
|
||||
|
||||
达梦、人大金仓、Oracle、DB2、Hive、Trino、Snowflake、SAP HANA 等 Agent/JDBC 数据库,需要匹配的 DBX Agent、JDBC 驱动和 JRE。请先在 DBX 中安装,再使用 MCP。
|
||||
Oracle、人大金仓和虚谷需要匹配的 DBX 原生 Agent,但不需要 JRE。达梦、DB2、Hive、Trino、Snowflake、SAP HANA 等 JDBC Agent 数据库需要匹配的 Agent、JDBC 驱动和 JRE。请先在 DBX 中安装所需组件,再使用 MCP。
|
||||
|
||||
### DBX Web / Docker 模式
|
||||
|
||||
|
|
|
|||
|
|
@ -105,7 +105,7 @@ PostgreSQL, MySQL, SQLite, compatible SQL databases, standalone Redis, and Mongo
|
|||
|
||||
### Agent/JDBC databases
|
||||
|
||||
Dameng, Kingbase, Oracle, DB2, Hive, Trino, Snowflake, SAP HANA, and other Agent/JDBC connections require the matching DBX Agent, JDBC driver, and JRE. Install them in DBX before using MCP.
|
||||
Oracle, KingBase, and XuguDB require their matching native DBX Agent but no JRE. Dameng, DB2, Hive, Trino, Snowflake, SAP HANA, and other JDBC Agent connections require the matching Agent, JDBC driver, and JRE. Install the required component in DBX before using MCP.
|
||||
|
||||
### DBX Web and Docker mode
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import assert from "node:assert/strict";
|
||||
import { afterEach, test, vi } from "vitest";
|
||||
import { buildAgentDownloadCatalog, downloadLinksFor, fetchAgentDownloadCatalog, formatSize } from "./agentRegistry";
|
||||
import { buildAgentDownloadCatalog, buildNativeAgentEntries, downloadLinksFor, fetchAgentDownloadCatalog, formatSize } from "./agentRegistry";
|
||||
|
||||
afterEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
|
|
@ -66,3 +66,31 @@ test("catalog falls back from GitHub to CNB", async () => {
|
|||
test("unknown fallback asset sizes render as unavailable", () => {
|
||||
assert.equal(formatSize(0), "—");
|
||||
});
|
||||
|
||||
test("KingBase release executables are listed as native agents", () => {
|
||||
const entries = buildNativeAgentEntries([
|
||||
{
|
||||
name: "dbx-agent-kingbase-windows-x64.exe",
|
||||
browser_download_url: "https://example.com/dbx-agent-kingbase-windows-x64.exe",
|
||||
size: 1024,
|
||||
},
|
||||
{
|
||||
name: "dbx-agent-kingbase-linux-x64",
|
||||
browser_download_url: "https://example.com/dbx-agent-kingbase-linux-x64",
|
||||
size: 2048,
|
||||
},
|
||||
{
|
||||
name: "dbx-agent-kingbase.jar",
|
||||
browser_download_url: "https://example.com/dbx-agent-kingbase.jar",
|
||||
size: 4096,
|
||||
},
|
||||
]);
|
||||
|
||||
assert.deepEqual(
|
||||
entries.map(({ key, platformKey, filename }) => ({ key, platformKey, filename })),
|
||||
[
|
||||
{ key: "kingbase", platformKey: "linux-x64", filename: "dbx-agent-kingbase-linux-x64" },
|
||||
{ key: "kingbase", platformKey: "windows-x64", filename: "dbx-agent-kingbase-windows-x64.exe" },
|
||||
],
|
||||
);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@ const GITHUB_RELEASE_DOWNLOAD_PREFIX = "https://github.com/t8y2/dbx/releases/dow
|
|||
const CNB_RELEASE_DOWNLOAD_PREFIX = "https://cnb.cool/dbxio.com/dbx/-/releases/download/";
|
||||
const MIN_APP_VERSION = "0.6.0";
|
||||
const driverVersionMap = driverVersions as Record<string, string>;
|
||||
const nativeDriverKeys = new Set(["oracle", "xugu"]);
|
||||
const nativeDriverKeys = new Set(["oracle", "kingbase", "xugu"]);
|
||||
|
||||
const platformLabels: Record<string, string> = {
|
||||
"macos-aarch64": "macOS (Apple Silicon)",
|
||||
|
|
@ -311,10 +311,11 @@ export function buildNativeAgentEntries(assets: GitHubReleaseAsset[]): NativeAge
|
|||
const entries: NativeAgentDisplayEntry[] = [];
|
||||
|
||||
for (const asset of assets) {
|
||||
const match = /^dbx-agent-(oracle|xugu)-(.+?)(?:\.exe)?$/.exec(asset.name);
|
||||
const match = /^dbx-agent-(.+?)-(macos-aarch64|macos-x64|linux-aarch64|linux-x64|windows-aarch64|windows-x64)(?:\.exe)?$/.exec(asset.name);
|
||||
if (!match) continue;
|
||||
|
||||
const [, key, platformKey] = match;
|
||||
if (!nativeDriverKeys.has(key)) continue;
|
||||
entries.push({
|
||||
key,
|
||||
label: labelForDriver(key),
|
||||
|
|
|
|||
Loading…
Reference in New Issue