feat(kingbase): add native Go agent

This commit is contained in:
t8y2 2026-07-20 18:43:23 +08:00
parent f5c614804a
commit 35c33a4085
20 changed files with 2886 additions and 44 deletions

View File

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

View File

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

View File

@ -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 agentDBX 会自动下载并管理 JRE 21 安装。
多数 Java agent 以 JRE 21 为目标。原生 agent`oracle`、`kingbase``xugu`)不需要 JRE。对 Java agentDBX 会自动下载并管理 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`
## 版本管理

View File

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

View File

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

View File

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

View File

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

View File

@ -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()", &current); 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
}

View File

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

View File

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

View File

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

View File

@ -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: "正在加载驱动列表...",

View File

@ -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/<原生模块>

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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