diff --git a/.github/scripts/bump-agent-versions.mjs b/.github/scripts/bump-agent-versions.mjs index 0cc6420d1c..6b6bb6f814 100644 --- a/.github/scripts/bump-agent-versions.mjs +++ b/.github/scripts/bump-agent-versions.mjs @@ -52,6 +52,7 @@ const nativeDriverDirectories = { cassandra: "cassandra-go", duckdb: "duckdb", hive: "hive-go", + argo: "argo-go", oracle: "oracle-go", kingbase: "kingbase-go", iotdb: "iotdb", @@ -67,13 +68,18 @@ const nativeDriverDirectories = { const crateNativeDriverDirectories = { "sqlite-worker": "crates/dbx-sqlite-worker", }; -const nativeDriverModules = new Set(["cassandra", "duckdb", "hive", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "rabbitmq", "rocketmq", "zookeeper", "tdengine", "etcd", "etcd2", "sqlite-worker"]); +const nativeDriverModules = new Set(["cassandra", "duckdb", "hive", "argo", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "rabbitmq", "rocketmq", "zookeeper", "tdengine", "etcd", "etcd2", "sqlite-worker"]); const nativeDriverSharedPaths = { hive: [ "agents/go-common/go-gssapi", "agents/go-common/gohive", "agents/go-common/gosasl", ], + argo: [ + "agents/go-common/go-gssapi", + "agents/go-common/gohive", + "agents/go-common/gosasl", + ], zookeeper: [ "agents/go-common/go-gssapi", "agents/go-common/gosasl", diff --git a/.github/workflows/agents-release.yml b/.github/workflows/agents-release.yml index 67a0b9067f..a6e7a4ec34 100644 --- a/.github/workflows/agents-release.yml +++ b/.github/workflows/agents-release.yml @@ -547,6 +547,46 @@ jobs: name: hive-native path: "release-native/dbx-agent-hive-*" + build-argo-native: + needs: [bump-versions] + if: ${{ contains(fromJSON(needs.bump-versions.outputs.native_modules), 'argo') }} + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.23.x" + - name: Test Argo native agent + working-directory: agents/drivers/argo-go + run: go test ./... + - name: Cross-compile Argo native agent + shell: bash + run: | + mkdir -p release-native + cd agents/drivers/argo-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-argo-${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: argo-native + path: "release-native/dbx-agent-argo-*" + build-kingbase-native: needs: [bump-versions] if: ${{ contains(fromJSON(needs.bump-versions.outputs.native_modules), 'kingbase') }} @@ -1217,7 +1257,7 @@ jobs: retention-days: 1 release: - needs: [bump-versions, commit-versions, build-agents, build-oracle-native, build-xugu-native, build-rabbitmq-native, build-rocketmq-native, build-etcd-native, build-etcd2-native, build-zookeeper-native, build-cassandra-native, build-hive-native, build-kingbase-native, build-vastbase-native, build-neo4j-native, build-iotdb-native, build-duckdb-native, build-sqlite-worker-native, build-tdengine-native, build-jre, reuse-previous-assets] + needs: [bump-versions, commit-versions, build-agents, build-oracle-native, build-xugu-native, build-rabbitmq-native, build-rocketmq-native, build-etcd-native, build-etcd2-native, build-zookeeper-native, build-cassandra-native, build-hive-native, build-argo-native, build-kingbase-native, build-vastbase-native, build-neo4j-native, build-iotdb-native, build-duckdb-native, build-sqlite-worker-native, build-tdengine-native, build-jre, reuse-previous-assets] if: ${{ always() && !contains(needs.*.result, 'failure') && !contains(needs.*.result, 'cancelled') }} runs-on: ubuntu-latest steps: @@ -1349,6 +1389,7 @@ jobs: zookeeper) echo "Apache ZooKeeper" ;; cassandra) echo "Apache Cassandra" ;; hive) echo "Apache Hive" ;; + argo) echo "星环Argo" ;; neo4j) echo "Neo4j" ;; iotdb) echo "Apache IoTDB" ;; tdengine) echo "TDengine" ;; @@ -1417,7 +1458,7 @@ jobs: [ -n "$DRIVERS" ] && DRIVERS="${DRIVERS},"$'\n' DRIVERS="${DRIVERS}$(generate_jar_entry "$name" "$label" "$f" "$jre_key" "$version" "$external_driver" "$native_json")" done - for name in oracle xugu kingbase vastbase neo4j iotdb duckdb sqlite-worker rabbitmq rocketmq zookeeper cassandra hive tdengine etcd etcd2; do + for name in oracle xugu kingbase vastbase neo4j iotdb duckdb sqlite-worker rabbitmq rocketmq zookeeper cassandra hive argo tdengine etcd etcd2; do version=$(get_module_version "$name") [ -f "release/dbx-agent-${name}-${version}.jar" ] && continue native_json=$(generate_native_platforms "$name" "$version") @@ -1505,6 +1546,7 @@ jobs: zookeeper) echo "Apache ZooKeeper" ;; cassandra) echo "Apache Cassandra" ;; hive) echo "Apache Hive" ;; + argo) echo "星环Argo" ;; neo4j) echo "Neo4j" ;; iotdb) echo "Apache IoTDB" ;; tdengine) echo "TDengine" ;; @@ -1554,6 +1596,8 @@ jobs: LOG_PATH="agents/drivers/cassandra-go/" elif [ "$name" = "hive" ]; then LOG_PATH="agents/drivers/hive-go/" + elif [ "$name" = "argo" ]; then + LOG_PATH="agents/drivers/argo-go/" elif [ "$name" = "neo4j" ]; then LOG_PATH="agents/drivers/neo4j-go/" elif [ "$name" = "iotdb" ]; then @@ -1572,6 +1616,12 @@ jobs: "agents/go-common/gohive/" "agents/go-common/gosasl/" ) + elif [ "$name" = "argo" ]; then + LOG_PATHS+=( + "agents/go-common/go-gssapi/" + "agents/go-common/gohive/" + "agents/go-common/gosasl/" + ) elif [ "$name" = "zookeeper" ]; then LOG_PATHS+=( "agents/go-common/go-gssapi/" diff --git a/agents/drivers/argo-go/MIGRATION_PARITY.md b/agents/drivers/argo-go/MIGRATION_PARITY.md new file mode 100644 index 0000000000..50cc38ef5e --- /dev/null +++ b/agents/drivers/argo-go/MIGRATION_PARITY.md @@ -0,0 +1,30 @@ +# Argo Go agent notes + +Status date: 2026-09-03. + +This agent is a fork of `agents/drivers/hive-go` created to serve 星环Argo +(Transwarp ArgoDB) connections exclusively: `supportsRoutines()` returns true +unconditionally and connection identity reports `ArgoDB (Transwarp)` / +`DBX ArgoDB Go Agent`. Vanilla Hive / Kyuubi / Impala stay on hive-go. + +No Hive 3 / Hive 4 / Kyuubi parity is claimed for this directory. The +validation matrix in `agents/drivers/hive-go/MIGRATION_PARITY.md` applies to +hive-go only. + +## Maintenance rule + +Most files are kept byte-identical to hive-go. Fixes landing in hive-go that +touch `metadata.go` or `main.go` must be ported here (adjusting only the +argo-specific branding, the unconditional `supportsRoutines()`, and log +prefixes), including their regression tests. + +## Validation + +Validated against Transwarp ArgoDB by the PR author (#7933): + +- `go test ./...` green in this module (routine views asserted queried + unconditionally; ArgoDB identity asserted) +- End-to-end on a local build with an ArgoDB connection: 星环Argo branding + shown, stored-procedure source opens as one whole statement, `CREATE OR + REPLACE PROCEDURE` executes as a single statement, and execute errors + surface readable diagnostics diff --git a/agents/drivers/argo-go/bench/README.md b/agents/drivers/argo-go/bench/README.md new file mode 100644 index 0000000000..9d4d84a000 --- /dev/null +++ b/agents/drivers/argo-go/bench/README.md @@ -0,0 +1,101 @@ +# Hive Agent benchmark + +This benchmark compares the same DBX JSON-RPC operations through the native +Go Hive Agent and the archived JDBC Hive Agent. Both candidates run on the same +host and connect to the same HiveServer2 instance. + +Use `functional_probe.py` before running performance benchmarks. It validates +connection, sessions, query values, metadata, pagination, failed-SQL semantics, +clean shutdown, and Java/Go result parity without concurrency load. + +The runner measures: + +- process startup and fresh connection latency; +- artifact size, idle RSS, and peak RSS; +- `SELECT 1`-shape lookup and 100/1,000/10,000-row decoding; +- `list_databases`, `list_tables`, and complete paged reads; +- 1, 8, and 32 concurrent DBX Agent sessions; +- mean, p50, p95, p99, throughput, and clean shutdown behavior. + +Candidate order rotates between rounds to reduce warm-cache and server-order +bias. Startup and connection samples always use a fresh Agent process. + +## Prepare the fixture + +The defaults expect `dbx_agent_bench.agent_bench` with exactly 10,000 or more +rows and columns named `id` and `payload`: + +```sql +CREATE DATABASE IF NOT EXISTS dbx_agent_bench; +CREATE TABLE IF NOT EXISTS dbx_agent_bench.agent_bench ( + id BIGINT, + payload STRING +) STORED AS ORC; +``` + +Populate deterministic rows before running the benchmark. Keep the fixture, +HiveServer2 configuration, Agent host, and Java runtime unchanged between +candidates. + +## Run + +Functional parity probe: + +```bash +GO_AGENT=/tmp/dbx-hive-bench/hive-agent-linux-amd64 \ +JDBC_AGENT_JAR=/tmp/dbx-hive-bench/dbx-agent-hive.jar \ +JAVA_BIN=/tmp/dbx-hive-bench/jre21/bin/java \ +HIVE_HOST=127.0.0.1 \ +HIVE_PORT=10000 \ +HIVE_DATABASE=dbx_agent_bench \ +HIVE_URL_PARAMS=auth=noSasl \ +python3 agents/drivers/hive-go/bench/functional_probe.py \ + > /tmp/dbx-hive-bench/functional-result.json +``` + +The functional probe runs the Go candidate only by default. Set +`BENCH_CANDIDATES=go,jdbc` only when an explicit historical-JDBC comparison is +needed. + +Performance benchmark: + +```bash +GO_AGENT=/tmp/dbx-hive-bench/hive-agent-linux-amd64 \ +JDBC_AGENT_JAR=/tmp/dbx-hive-bench/dbx-agent-hive.jar \ +HIVE_HOST=127.0.0.1 \ +HIVE_PORT=10000 \ +HIVE_DATABASE=dbx_agent_bench \ +HIVE_URL_PARAMS=auth=noSasl \ +python3 agents/drivers/hive-go/bench/agent_compare.py \ + > /tmp/dbx-hive-bench/result.json +``` + +## Configuration + +- `BENCH_CANDIDATES`: performance benchmark default `go,jdbc`; functional probe + default `go`. +- `GO_AGENT_COMMAND`, `JDBC_AGENT_COMMAND`: optional full launch commands. +- `BENCH_STARTUPS`, `BENCH_CONNECTS`: fresh-process sample counts, default `8`. +- `BENCH_ROUNDS`: alternating steady-state rounds, default `3`. +- `BENCH_WARMUPS`: warmups before each workload, default `2`. +- `BENCH_CONCURRENCY`: comma-separated session counts, default `1,8,32`. +- `BENCH_CONCURRENCY_OPS_PER_WORKER`: operations per session, default `8`. +- `HIVE_HOST`, `HIVE_PORT`, `HIVE_DATABASE`, `HIVE_BENCH_TABLE`. +- `HIVE_USERNAME`, `HIVE_PASSWORD`, `HIVE_URL_PARAMS`, `HIVE_CONNECTION_STRING`. +- `HIVE_SSL`, `HIVE_CA_CERT_PATH`, `HIVE_CLIENT_CERT_PATH`, `HIVE_CLIENT_KEY_PATH`. +- `BENCH_*_SQL` and `BENCH_*_COUNT` override individual workloads. + +Run the benchmark on the Agent host. Do not compare a local Go process with a +remote JDBC process, use different HS2 endpoints, or mutate the fixture between +candidates. + +## Kerberos fixture + +`kdc_fixture` starts a test-only in-process KDC and writes a temporary +`krb5.conf` and keytab containing `alice` and `hive/localhost`. Never use these +credentials outside an isolated compatibility environment, and delete the +generated directory after validation. + +```bash +go run ./bench/kdc_fixture -dir /tmp/dbx-hive-kerberos +``` diff --git a/agents/drivers/argo-go/bench/agent_compare.py b/agents/drivers/argo-go/bench/agent_compare.py new file mode 100755 index 0000000000..f2c68d8994 --- /dev/null +++ b/agents/drivers/argo-go/bench/agent_compare.py @@ -0,0 +1,664 @@ +#!/usr/bin/env python3 +import json +import os +import queue +import shlex +import statistics +import subprocess +import sys +import threading +import time +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path + + +@dataclass(frozen=True) +class Candidate: + name: str + command: list[str] + artifact: Path + rss_command: str = "" + + +class AgentProcess: + def __init__(self, candidate: Candidate): + self.candidate = candidate + self.process = subprocess.Popen( + candidate.command, + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + bufsize=1, + ) + self.request_id = 0 + self.request_lock = threading.Lock() + self.write_lock = threading.Lock() + self.pending: dict[int, queue.Queue] = {} + self.ready = threading.Event() + self.saw_ready = False + self.exited = threading.Event() + self.stderr_lines: list[str] = [] + threading.Thread(target=self._read_stdout, daemon=True).start() + threading.Thread(target=self._drain_stderr, daemon=True).start() + if not self.ready.wait(env_float("BENCH_READY_TIMEOUT", 30.0)) or not self.saw_ready: + raise TimeoutError(self._failure("timed out waiting for agent readiness")) + + def _read_stdout(self) -> None: + assert self.process.stdout is not None + for line in self.process.stdout: + try: + response = json.loads(line) + except json.JSONDecodeError: + continue + if response.get("ready") is True: + self.saw_ready = True + self.ready.set() + continue + response_id = response.get("id") + if not isinstance(response_id, int): + continue + with self.request_lock: + response_queue = self.pending.get(response_id) + if response_queue is not None: + response_queue.put(response) + self.exited.set() + self.ready.set() + with self.request_lock: + pending = list(self.pending.values()) + for response_queue in pending: + response_queue.put(RuntimeError(self._failure("agent process exited"))) + + def _drain_stderr(self) -> None: + assert self.process.stderr is not None + for line in self.process.stderr: + self.stderr_lines.append(line.rstrip()) + + def call(self, method: str, params: dict | None = None) -> object: + if self.exited.is_set(): + raise RuntimeError(self._failure(f"agent exited before {method}")) + with self.request_lock: + self.request_id += 1 + request_id = self.request_id + response_queue: queue.Queue = queue.Queue(maxsize=1) + self.pending[request_id] = response_queue + request = { + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params or {}, + } + try: + assert self.process.stdin is not None + with self.write_lock: + self.process.stdin.write(json.dumps(request, separators=(",", ":")) + "\n") + self.process.stdin.flush() + response = response_queue.get(timeout=env_float("BENCH_RPC_TIMEOUT", 180.0)) + except queue.Empty as error: + raise TimeoutError(self._failure(f"timed out during {method}")) from error + finally: + with self.request_lock: + self.pending.pop(request_id, None) + if isinstance(response, Exception): + raise response + if response.get("error") is not None: + raise RuntimeError( + f"{self.candidate.name} {method}: " + f"{json.dumps(response['error'], ensure_ascii=False)}" + ) + return response.get("result") + + def rss_kib(self) -> int: + if self.candidate.rss_command: + output = subprocess.check_output( + self.candidate.rss_command, + shell=True, + text=True, + ).strip() + return int(output or "0") + status_path = Path(f"/proc/{self.process.pid}/status") + if status_path.is_file(): + for line in status_path.read_text().splitlines(): + if line.startswith("VmRSS:"): + return int(line.split()[1]) + output = subprocess.check_output( + ["ps", "-o", "rss=", "-p", str(self.process.pid)], + text=True, + ).strip() + return int(output or "0") + + def close(self) -> bool: + if self.process.poll() is not None: + return True + try: + self.call("shutdown") + except Exception: + pass + try: + self.process.wait(timeout=3) + return True + except subprocess.TimeoutExpired: + self.process.terminate() + try: + self.process.wait(timeout=2) + except subprocess.TimeoutExpired: + self.process.kill() + self.process.wait(timeout=5) + return False + + def _failure(self, message: str) -> str: + stderr = "\n".join(self.stderr_lines[-20:]) + return f"{self.candidate.name}: {message}\n{stderr}".rstrip() + + +class RSSMonitor: + def __init__(self, process: AgentProcess): + self.process = process + self.peak_kib = 0 + self.stop_event = threading.Event() + self.thread = threading.Thread(target=self._run, daemon=True) + + def __enter__(self): + self.peak_kib = self.process.rss_kib() + self.thread.start() + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.stop_event.set() + self.thread.join(timeout=1) + try: + self.peak_kib = max(self.peak_kib, self.process.rss_kib()) + except (OSError, subprocess.SubprocessError, ValueError): + pass + + def _run(self) -> None: + interval = env_float("BENCH_RSS_INTERVAL", 0.02) + while not self.stop_event.wait(interval): + try: + self.peak_kib = max(self.peak_kib, self.process.rss_kib()) + except (OSError, subprocess.SubprocessError, ValueError): + return + + +def main() -> None: + candidates = configured_candidates() + connection = connection_params() + startup_iterations = env_int("BENCH_STARTUPS", 8) + connect_iterations = env_int("BENCH_CONNECTS", 8) + rounds = env_int("BENCH_ROUNDS", 3) + warmups = env_int("BENCH_WARMUPS", 2) + workloads = configured_workloads(connection["database"]) + concurrency_levels = env_int_list("BENCH_CONCURRENCY", [1, 8, 32]) + + startup_samples = {candidate.name: [] for candidate in candidates} + connect_samples = {candidate.name: [] for candidate in candidates} + workload_samples = { + candidate.name: {workload["name"]: [] for workload in workloads} + for candidate in candidates + } + concurrency_samples = { + candidate.name: {str(level): [] for level in concurrency_levels} + for candidate in candidates + } + process_metrics = {candidate.name: [] for candidate in candidates} + + for iteration in range(startup_iterations): + for candidate in rotated(candidates, iteration): + startup_samples[candidate.name].append(benchmark_startup(candidate)) + + for iteration in range(connect_iterations): + for candidate in rotated(candidates, iteration): + connect_samples[candidate.name].append(benchmark_connect(candidate, connection)) + + for round_index in range(rounds): + for candidate in rotated(candidates, round_index): + process = AgentProcess(candidate) + shutdown_clean = False + metrics = None + try: + process.call("connect", connection) + idle_rss_kib = process.rss_kib() + with RSSMonitor(process) as monitor: + for workload in workloads: + sample = benchmark_workload(process, workload, warmups) + workload_samples[candidate.name][workload["name"]].append(sample) + process.call("disconnect") + metrics = { + "idle_rss_kib": idle_rss_kib, + "peak_rss_kib": monitor.peak_kib, + } + finally: + shutdown_clean = process.close() + if metrics is not None: + metrics["shutdown_exited_within_3s"] = shutdown_clean + process_metrics[candidate.name].append(metrics) + + for level in concurrency_levels: + for candidate in rotated(candidates, round_index + level): + concurrency_samples[candidate.name][str(level)].append( + benchmark_concurrency(candidate, connection, level, warmups) + ) + + results = [] + for candidate in candidates: + name = candidate.name + results.append( + { + "candidate": name, + "command": candidate.command, + "artifact_bytes": candidate.artifact.stat().st_size, + "startup": summarize_latencies(startup_samples[name]), + "connect": summarize_latencies(connect_samples[name]), + "process": summarize_process_metrics(process_metrics[name]), + "workloads": [ + summarize_rounds(workload["name"], workload_samples[name][workload["name"]]) + for workload in workloads + ], + "concurrency": [ + summarize_rounds(f"concurrency_{level}", concurrency_samples[name][str(level)]) + for level in concurrency_levels + ], + } + ) + + output = { + "host": os.uname().nodename, + "server": env_default("HIVE_SERVER", f"{connection['host']}:{connection['port']}"), + "database": connection["database"], + "table": env_default("HIVE_BENCH_TABLE", "agent_bench"), + "startup_iterations": startup_iterations, + "connect_iterations": connect_iterations, + "rounds": rounds, + "warmups": warmups, + "concurrency_levels": concurrency_levels, + "results": results, + } + json.dump(output, sys.stdout, ensure_ascii=False, indent=2) + sys.stdout.write("\n") + + +def configured_candidates() -> list[Candidate]: + selected = { + item.strip() + for item in env_default("BENCH_CANDIDATES", "go,jdbc").split(",") + if item.strip() + } + candidates = [] + if "go" in selected: + artifact = required_path("GO_AGENT") + raw_command = os.getenv("GO_AGENT_COMMAND", "") + command = shlex.split(raw_command) if raw_command else [str(artifact)] + candidates.append( + Candidate("go-native", command, artifact, os.getenv("GO_RSS_COMMAND", "")) + ) + if "jdbc" in selected: + artifact = required_path("JDBC_AGENT_JAR") + raw_command = os.getenv("JDBC_AGENT_COMMAND", "") + command = ( + shlex.split(raw_command) + if raw_command + else [env_default("JAVA_BIN", "java"), "-jar", str(artifact)] + ) + candidates.append( + Candidate("jdbc-java", command, artifact, os.getenv("JDBC_RSS_COMMAND", "")) + ) + if not candidates: + raise ValueError("BENCH_CANDIDATES selected no candidates") + return candidates + + +def connection_params() -> dict: + return { + "host": env_default("HIVE_HOST", "127.0.0.1"), + "port": env_int("HIVE_PORT", 10000), + "database": env_default("HIVE_DATABASE", "dbx_agent_bench"), + "username": os.getenv("HIVE_USERNAME", ""), + "password": os.getenv("HIVE_PASSWORD", ""), + "url_params": env_default("HIVE_URL_PARAMS", "auth=noSasl"), + "connection_string": os.getenv("HIVE_CONNECTION_STRING", ""), + "ssl": env_bool("HIVE_SSL", False), + "ca_cert_path": os.getenv("HIVE_CA_CERT_PATH", ""), + "client_cert_path": os.getenv("HIVE_CLIENT_CERT_PATH", ""), + "client_key_path": os.getenv("HIVE_CLIENT_KEY_PATH", ""), + "connect_timeout_secs": env_int("HIVE_CONNECT_TIMEOUT", 30), + } + + +def configured_workloads(database: str) -> list[dict]: + table = env_default("HIVE_BENCH_TABLE", "agent_bench") + qualified = f"`{database}`.`{table}`" + workloads = [ + query_workload("select_one", "SELECT 1 AS value", 1, 40), + query_workload("rows_100", f"SELECT id, payload FROM {qualified} LIMIT 100", 100, 20), + query_workload("rows_1000", f"SELECT id, payload FROM {qualified} LIMIT 1000", 1000, 10), + query_workload("rows_10000", f"SELECT id, payload FROM {qualified} LIMIT 10000", 10000, 3), + { + "name": "list_databases", + "kind": "rpc", + "method": "list_databases", + "params": {}, + "count": env_int("BENCH_LIST_DATABASES_COUNT", 20), + }, + { + "name": "list_tables", + "kind": "rpc", + "method": "list_tables", + "params": {"schema": database}, + "count": env_int("BENCH_LIST_TABLES_COUNT", 20), + }, + { + "name": "page_10000_by_500", + "kind": "paged", + "sql": env_default("BENCH_PAGE_SQL", f"SELECT id, payload FROM {qualified} LIMIT 10000"), + "max_rows": 10000, + "page_size": env_int("BENCH_PAGE_SIZE", 500), + "count": env_int("BENCH_PAGE_COUNT", 3), + }, + ] + return workloads + + +def query_workload(name: str, fallback_sql: str, max_rows: int, fallback_count: int) -> dict: + suffix = name.upper() + return { + "name": name, + "kind": "rpc", + "method": "execute_query", + "params": { + "sql": env_default(f"BENCH_{suffix}_SQL", fallback_sql), + "maxRows": max_rows, + "fetchSize": min(max_rows, env_int("BENCH_FETCH_SIZE", 1000)), + }, + "count": env_int(f"BENCH_{suffix}_COUNT", fallback_count), + } + + +def benchmark_startup(candidate: Candidate) -> float: + started = time.perf_counter() + process = AgentProcess(candidate) + elapsed = elapsed_ms(started) + process.close() + return elapsed + + +def benchmark_connect(candidate: Candidate, connection: dict) -> float: + process = AgentProcess(candidate) + try: + started = time.perf_counter() + process.call("connect", connection) + return elapsed_ms(started) + finally: + process.close() + + +def benchmark_workload(process: AgentProcess, workload: dict, warmups: int) -> dict: + for _ in range(warmups): + execute_workload(process, workload) + samples = [] + started = time.perf_counter() + for _ in range(workload["count"]): + operation_started = time.perf_counter() + execute_workload(process, workload) + samples.append(elapsed_ms(operation_started)) + elapsed = time.perf_counter() - started + return sample_result(workload["count"], elapsed, samples) + + +def benchmark_concurrency( + candidate: Candidate, + connection: dict, + concurrency: int, + warmups: int, +) -> dict: + process = AgentProcess(candidate) + session_ids = [f"bench-{concurrency}-{index}" for index in range(concurrency)] + qualified = ( + f"`{connection['database']}`." + f"`{env_default('HIVE_BENCH_TABLE', 'agent_bench')}`" + ) + workload = { + "kind": "rpc", + "method": "execute_query", + "params": { + "sql": env_default( + "BENCH_CONCURRENCY_SQL", + "SELECT 1 AS value", + ), + "maxRows": 1, + }, + } + operations_per_worker = env_int("BENCH_CONCURRENCY_OPS_PER_WORKER", 8) + try: + for session_id in session_ids: + process.call("open_session", {**connection, "agentSessionId": session_id}) + for session_id in session_ids: + for _ in range(warmups): + execute_workload(process, workload, session_id) + with RSSMonitor(process) as monitor: + started = time.perf_counter() + with ThreadPoolExecutor(max_workers=concurrency) as executor: + futures = [ + executor.submit( + concurrency_worker, + process, + workload, + session_id, + operations_per_worker, + ) + for session_id in session_ids + ] + samples = [sample for future in futures for sample in future.result()] + elapsed = time.perf_counter() - started + result = sample_result(len(samples), elapsed, samples) + result["concurrency"] = concurrency + result["peak_rss_kib"] = monitor.peak_kib + return result + finally: + for session_id in session_ids: + try: + process.call("close_session", {"agentSessionId": session_id}) + except Exception: + pass + process.close() + + +def concurrency_worker( + process: AgentProcess, + workload: dict, + session_id: str, + operations: int, +) -> list[float]: + samples = [] + for _ in range(operations): + started = time.perf_counter() + execute_workload(process, workload, session_id) + samples.append(elapsed_ms(started)) + return samples + + +def execute_workload( + process: AgentProcess, + workload: dict, + agent_session_id: str = "", +) -> object: + params = dict(workload.get("params", {})) + if agent_session_id: + params["agentSessionId"] = agent_session_id + if workload["kind"] == "rpc": + return process.call(workload["method"], params) + if workload["kind"] != "paged": + raise ValueError(f"unknown workload kind: {workload['kind']}") + first = process.call( + "execute_query_page", + { + "sql": workload["sql"], + "maxRows": workload["max_rows"], + "pageSize": workload["page_size"], + **({"agentSessionId": agent_session_id} if agent_session_id else {}), + }, + ) + rows = len(first.get("rows", [])) + session_id = first.get("session_id") + has_more = first.get("has_more", False) + try: + while has_more: + page = process.call( + "fetch_query_page", + { + "sessionId": session_id, + "pageSize": workload["page_size"], + **({"agentSessionId": agent_session_id} if agent_session_id else {}), + }, + ) + rows += len(page.get("rows", [])) + session_id = page.get("session_id") + has_more = page.get("has_more", False) + finally: + if session_id: + process.call( + "close_query_session", + { + "sessionId": session_id, + **({"agentSessionId": agent_session_id} if agent_session_id else {}), + }, + ) + if rows != workload["max_rows"]: + raise RuntimeError( + f"{process.candidate.name} paged query returned {rows} rows, " + f"expected {workload['max_rows']}" + ) + return rows + + +def sample_result(count: int, elapsed: float, samples: list[float]) -> dict: + summary = summarize_latencies(samples) + summary.update( + { + "count": count, + "elapsed_ms": elapsed * 1000, + "ops_per_sec": count / elapsed, + } + ) + return summary + + +def summarize_latencies(samples: list[float]) -> dict: + ordered = sorted(samples) + return { + "samples_ms": samples, + "mean_ms": statistics.mean(samples), + "p50_ms": percentile(ordered, 0.50), + "p95_ms": percentile(ordered, 0.95), + "p99_ms": percentile(ordered, 0.99), + "min_ms": ordered[0], + "max_ms": ordered[-1], + } + + +def summarize_rounds(name: str, rounds: list[dict]) -> dict: + latencies = [sample for round_result in rounds for sample in round_result["samples_ms"]] + elapsed = sum(round_result["elapsed_ms"] for round_result in rounds) / 1000 + result = summarize_latencies(latencies) + result.update( + { + "name": name, + "rounds": rounds, + "count": len(latencies), + "elapsed_ms": elapsed * 1000, + "ops_per_sec": len(latencies) / elapsed, + } + ) + peak_values = [round_result.get("peak_rss_kib", 0) for round_result in rounds] + if any(peak_values): + result["peak_rss_kib"] = max(peak_values) + return result + + +def summarize_process_metrics(samples: list[dict]) -> dict: + return { + "idle_rss_kib": summarize_numbers([sample["idle_rss_kib"] for sample in samples]), + "peak_rss_kib": summarize_numbers([sample["peak_rss_kib"] for sample in samples]), + "shutdown_exited_within_3s": all( + sample["shutdown_exited_within_3s"] for sample in samples + ), + "rounds": samples, + } + + +def summarize_numbers(values: list[int]) -> dict: + return { + "min": min(values), + "median": statistics.median(values), + "max": max(values), + } + + +def rotated(values: list[Candidate], offset: int) -> list[Candidate]: + if not values: + return [] + shift = offset % len(values) + return values[shift:] + values[:shift] + + +def percentile(values: list[float], fraction: float) -> float: + if not values: + return 0.0 + index = min(len(values) - 1, max(0, round((len(values) - 1) * fraction))) + return values[index] + + +def elapsed_ms(started: float) -> float: + return (time.perf_counter() - started) * 1000 + + +def required_path(name: str) -> Path: + value = os.getenv(name, "") + if not value: + raise ValueError(f"{name} is required") + path = Path(value).expanduser().resolve() + if not path.is_file(): + raise FileNotFoundError(path) + return path + + +def env_default(name: str, fallback: str) -> str: + return os.getenv(name, "") or fallback + + +def env_int(name: str, fallback: int) -> int: + value = int(env_default(name, str(fallback))) + if value < 1: + raise ValueError(f"{name} must be positive") + return value + + +def env_int_list(name: str, fallback: list[int]) -> list[int]: + raw = os.getenv(name, "") + values = fallback if not raw else [int(value.strip()) for value in raw.split(",")] + if not values or any(value < 1 for value in values): + raise ValueError(f"{name} must contain positive integers") + return values + + +def env_float(name: str, fallback: float) -> float: + value = float(env_default(name, str(fallback))) + if value <= 0: + raise ValueError(f"{name} must be positive") + return value + + +def env_bool(name: str, fallback: bool) -> bool: + raw = os.getenv(name) + if raw is None or raw == "": + return fallback + normalized = raw.strip().lower() + if normalized in {"1", "true", "yes", "on"}: + return True + if normalized in {"0", "false", "no", "off"}: + return False + raise ValueError(f"{name} must be a boolean") + + +if __name__ == "__main__": + main() diff --git a/agents/drivers/argo-go/bench/functional_probe.py b/agents/drivers/argo-go/bench/functional_probe.py new file mode 100644 index 0000000000..66562a0946 --- /dev/null +++ b/agents/drivers/argo-go/bench/functional_probe.py @@ -0,0 +1,203 @@ +#!/usr/bin/env python3 +import hashlib +import json +import os +import sys +from pathlib import Path + +from agent_compare import AgentProcess, configured_candidates, connection_params, env_default, env_int + + +def main() -> None: + os.environ.setdefault("BENCH_CANDIDATES", "go") + connection = connection_params() + candidates = configured_candidates() + results = {candidate.name: probe_candidate(candidate, connection) for candidate in candidates} + output = { + "server": env_default("HIVE_SERVER", f"{connection['host']}:{connection['port']}"), + "connection": sanitized_connection(connection), + "artifacts": { + candidate.name: { + "path": str(candidate.artifact), + "sha256": sha256(candidate.artifact), + "size_bytes": candidate.artifact.stat().st_size, + } + for candidate in candidates + }, + "results": results, + "parity": compare_results(results), + } + json.dump(output, sys.stdout, ensure_ascii=False, indent=2) + sys.stdout.write("\n") + if any(not result.get("ok") for result in results.values()) or not output["parity"]["ok"]: + raise SystemExit(1) + + +def probe_candidate(candidate, connection: dict) -> dict: + process = None + session_id = f"functional-probe-{candidate.name}" + result = {"ok": False} + try: + process = AgentProcess(candidate) + result["test_connection"] = process.call("test_connection", connection) + process.call("open_session", {"agentSessionId": session_id, **connection}) + result["validate_session"] = process.call("validate_session", {"agentSessionId": session_id}) + result["select_one"] = normalized_query( + process.call( + "execute_query", + { + "agentSessionId": session_id, + "sql": env_default("PROBE_SELECT_SQL", "SELECT 1 AS value"), + "maxRows": env_int("PROBE_SELECT_MAX_ROWS", 10), + "fetchSize": env_int("PROBE_FETCH_SIZE", 10), + }, + ) + ) + result["databases"] = sorted( + item.get("name", "") + for item in process.call("list_databases", {"agentSessionId": session_id}) + ) + schema = env_default("PROBE_SCHEMA", connection["database"]) + result["tables"] = sorted( + item.get("name", "") + for item in process.call( + "list_tables", + {"agentSessionId": session_id, "schema": schema}, + ) + ) + result["paging"] = probe_paging(process, session_id) + result["invalid_sql"] = probe_failure_semantics(process, session_id) + result["after_failure"] = normalized_query( + process.call( + "execute_query", + { + "agentSessionId": session_id, + "sql": env_default("PROBE_AFTER_FAILURE_SQL", "SELECT 2 AS value"), + "maxRows": 10, + "fetchSize": env_int("PROBE_FETCH_SIZE", 10), + }, + ) + ) + process.call("close_session", {"agentSessionId": session_id}) + result["ok"] = True + except Exception as error: + result["error"] = str(error) + finally: + if process is not None: + result["clean_shutdown"] = process.close() + return result + + +def probe_paging(process: AgentProcess, agent_session_id: str) -> dict: + page_size = env_int("PROBE_PAGE_SIZE", 2) + first = process.call( + "execute_query_page", + { + "agentSessionId": agent_session_id, + "sql": env_default( + "PROBE_PAGE_SQL", + "SELECT id, payload FROM dbx_agent_bench.agent_bench LIMIT 3", + ), + "maxRows": env_int("PROBE_PAGE_MAX_ROWS", 3), + "pageSize": page_size, + }, + ) + pages = [first] + query_session_id = first.get("session_id") + while pages[-1].get("has_more"): + pages.append( + process.call( + "fetch_query_page", + { + "agentSessionId": agent_session_id, + "sessionId": query_session_id, + "pageSize": page_size, + }, + ) + ) + return { + "columns": first.get("columns", []), + "column_types": first.get("column_types", []), + "rows": [row for page in pages for row in page.get("rows", [])], + "page_count": len(pages), + "has_more_final": pages[-1].get("has_more", False), + "truncated": any(page.get("truncated", False) for page in pages), + } + + +def probe_failure_semantics(process: AgentProcess, agent_session_id: str) -> dict: + try: + process.call( + "execute_query", + { + "agentSessionId": agent_session_id, + "sql": env_default( + "PROBE_INVALID_SQL", + "SELECT * FROM dbx_missing_table_for_failure_semantics", + ), + "maxRows": 10, + "fetchSize": env_int("PROBE_FETCH_SIZE", 10), + }, + ) + except Exception as error: + return {"failed": True, "error": str(error)} + return {"failed": False, "error": ""} + + +def normalized_query(result: dict) -> dict: + return { + "columns": result.get("columns", []), + "column_types": result.get("column_types", []), + "rows": result.get("rows", []), + "truncated": result.get("truncated", False), + } + + +def compare_results(results: dict) -> dict: + successful = [result for result in results.values() if result.get("ok")] + if len(successful) < 2: + return {"ok": len(results) == 1 and len(successful) == 1, "differences": []} + baseline = successful[0] + differences = [] + for field in ["select_one", "databases", "tables", "paging", "after_failure"]: + expected = baseline.get(field) + for candidate_name, candidate_result in results.items(): + if candidate_result.get("ok") and candidate_result.get(field) != expected: + differences.append( + { + "candidate": candidate_name, + "field": field, + "expected": expected, + "actual": candidate_result.get(field), + } + ) + for candidate_name, candidate_result in results.items(): + if candidate_result.get("ok") and not candidate_result.get("invalid_sql", {}).get("failed"): + differences.append( + { + "candidate": candidate_name, + "field": "invalid_sql.failed", + "expected": True, + "actual": False, + } + ) + return {"ok": not differences, "differences": differences} + + +def sanitized_connection(connection: dict) -> dict: + result = dict(connection) + if result.get("password"): + result["password"] = "***" + return result + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for chunk in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +if __name__ == "__main__": + main() diff --git a/agents/drivers/argo-go/bench/kdc_fixture/main.go b/agents/drivers/argo-go/bench/kdc_fixture/main.go new file mode 100644 index 0000000000..667d0eecca --- /dev/null +++ b/agents/drivers/argo-go/bench/kdc_fixture/main.go @@ -0,0 +1,99 @@ +package main + +import ( + "encoding/json" + "flag" + "fmt" + "log" + "os" + "os/signal" + "path/filepath" + "syscall" + + "github.com/jcmturner/krb5test" +) + +type fixtureInfo struct { + Realm string `json:"realm"` + Address string `json:"address"` + ConfigPath string `json:"config_path"` + KeytabPath string `json:"keytab_path"` + ClientPrincipal string `json:"client_principal"` + ServicePrincipal string `json:"service_principal"` + ZooKeeperPrincipal string `json:"zookeeper_principal"` +} + +func main() { + directory := flag.String("dir", "", "directory for generated Kerberos fixture files") + flag.Parse() + if *directory == "" { + log.Fatal("-dir is required") + } + if err := os.MkdirAll(*directory, 0o700); err != nil { + log.Fatal(err) + } + + logger := log.New(os.Stderr, "kdc: ", log.LstdFlags) + kdc, err := krb5test.NewKDC(map[string][]string{ + "alice": nil, + "hive/localhost": nil, + "zookeeper/localhost": nil, + }, logger) + if err != nil { + log.Fatal(err) + } + kdc.KRB5Conf.LibDefaults.UDPPreferenceLimit = 1 + kdc.Start() + defer kdc.Close() + + configPath := filepath.Join(*directory, "krb5.conf") + keytabPath := filepath.Join(*directory, "fixture.keytab") + config := fmt.Sprintf(`[libdefaults] + default_realm = %s + dns_lookup_realm = false + dns_lookup_kdc = false + rdns = false + udp_preference_limit = 1 + default_tgs_enctypes = aes256-cts-hmac-sha1-96 + default_tkt_enctypes = aes256-cts-hmac-sha1-96 + permitted_enctypes = aes256-cts-hmac-sha1-96 + +[realms] + %s = { + kdc = %s + } + +[domain_realm] + .localhost = %s + localhost = %s +`, kdc.Realm, kdc.Realm, kdc.TCPListener.Addr().String(), kdc.Realm, kdc.Realm) + if err := os.WriteFile(configPath, []byte(config), 0o644); err != nil { + log.Fatal(err) + } + keytab, err := kdc.Keytab.Marshal() + if err != nil { + log.Fatal(err) + } + if err := os.WriteFile(keytabPath, keytab, 0o600); err != nil { + log.Fatal(err) + } + info := fixtureInfo{ + Realm: kdc.Realm, + Address: kdc.TCPListener.Addr().String(), + ConfigPath: configPath, + KeytabPath: keytabPath, + ClientPrincipal: "alice@" + kdc.Realm, + ServicePrincipal: "hive/localhost@" + kdc.Realm, + ZooKeeperPrincipal: "zookeeper/localhost@" + kdc.Realm, + } + if err := json.NewEncoder(os.Stdout).Encode(info); err != nil { + log.Fatal(err) + } + if err := os.Stdout.Sync(); err != nil { + log.Fatal(err) + } + + signals := make(chan os.Signal, 1) + signal.Notify(signals, syscall.SIGINT, syscall.SIGTERM) + <-signals +} diff --git a/agents/drivers/argo-go/config.go b/agents/drivers/argo-go/config.go new file mode 100644 index 0000000000..7bc34e0ba7 --- /dev/null +++ b/agents/drivers/argo-go/config.go @@ -0,0 +1,1236 @@ +package main + +import ( + "bufio" + "crypto/tls" + "crypto/x509" + "encoding/base64" + "errors" + "fmt" + "io" + "net" + "net/url" + "os" + "path/filepath" + "regexp" + "runtime" + "strconv" + "strings" + "time" +) + +const ( + defaultHivePort = 10000 + defaultHiveDatabase = "default" + defaultHiveHTTPPath = "cliservice" + defaultHiveService = "hive" + defaultImpalaService = "impala" + defaultZooKeeperNamespace = "hiveserver2" + resultSetUniqueColumnNames = "hive.resultset.use.unique.column.names" + defaultConnectTimeout = 15 * time.Second + defaultRetryInterval = time.Second + defaultBrowserSSOTimeout = 120 * time.Second + defaultCookieName = "hive.server2.auth" +) + +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"` + SSL bool `json:"ssl"` + CACertPath string `json:"ca_cert_path"` + ClientCertPath string `json:"client_cert_path"` + ClientKeyPath string `json:"client_key_path"` + ConnectTimeout int `json:"connect_timeout_secs"` + AgentJavaOptions []string `json:"agent_java_options"` + SessionRole string `json:"sessionRole"` + DatabaseType string `json:"database_type"` +} + +type endpoint struct { + Host string + Port int + TransportMode string + HTTPPath string + Auth string + Principal string + SSL bool +} + +func (value endpoint) address() string { + return net.JoinHostPort(value.Host, strconv.Itoa(value.Port)) +} + +type kerberosConfig struct { + Enabled bool + ServerPrincipal string + ServerPrincipalExplicit bool + ClientPrincipal string + Service string + ServerName string + Realm string + ConfigPath string + JAASConfigPath string + KeytabPath string + CCachePath string + Password string + AuthorizationID string + QOP string + UseKeytab bool + UseTicketCache bool + UseSSPI bool + CanonicalHostname bool + ChannelBinding bool + DisablePAFXFAST bool +} + +type zooKeeperKerberosConfig struct { + Enabled bool + Service string + ServerPrincipal string + Realm string + CanonicalHostname bool +} + +type connectionConfig struct { + DatabaseType string + Endpoints []endpoint + Database string + Username string + Password string + Auth string + TransportMode string + HTTPPath string + TLSConfig *tls.Config + AuthExplicit bool + TransportModeExplicit bool + HTTPPathExplicit bool + TLSExplicit bool + HiveConfiguration map[string]string + ServiceDiscoveryMode string + ZooKeeperNamespace string + ZooKeeperAuthScheme string + ZooKeeperAuth string + ZooKeeperTLSConfig *tls.Config + ZooKeeperKerberos zooKeeperKerberosConfig + ConnectTimeout time.Duration + SocketTimeout time.Duration + FetchSize int + MaxMessageSize int32 + Retries int + RetryInterval time.Duration + HTTPHeaders map[string]string + HTTPCookies map[string]string + RequestTracking bool + CookieAuth bool + CookieName string + JWT string + DelegationToken string + BrowserToken string + BrowserClientID string + BrowserResponsePort int + BrowserResponseTimeout time.Duration + BrowserDisableSSLCheck bool + InitStatements []string + Kerberos kerberosConfig +} + +func parseConnectionConfig(params connectParams) (connectionConfig, error) { + hasStructuredEndpoint := strings.TrimSpace(params.Host) != "" + config := connectionConfig{ + DatabaseType: strings.ToLower(strings.TrimSpace(params.DatabaseType)), + Database: strings.TrimSpace(params.Database), + Username: params.Username, + Password: params.Password, + Auth: "NONE", + TransportMode: "binary", + HTTPPath: defaultHiveHTTPPath, + HiveConfiguration: map[string]string{}, + FetchSize: defaultFetchSize, + Retries: 1, + RetryInterval: defaultRetryInterval, + BrowserResponseTimeout: defaultBrowserSSOTimeout, + HTTPHeaders: map[string]string{}, + HTTPCookies: map[string]string{}, + CookieAuth: true, + CookieName: defaultCookieName, + ZooKeeperNamespace: defaultZooKeeperNamespace, + ZooKeeperKerberos: zooKeeperKerberosConfig{ + Service: "zookeeper", + CanonicalHostname: true, + }, + ConnectTimeout: defaultConnectTimeout, + Kerberos: kerberosConfig{ + Service: defaultHiveService, + QOP: "auth", + CanonicalHostname: true, + }, + } + if strings.EqualFold(params.DatabaseType, "impala") { + config.Auth = "NOSASL" + config.Kerberos.Service = defaultImpalaService + } + if params.ConnectTimeout > 0 { + config.ConnectTimeout = time.Duration(params.ConnectTimeout) * time.Second + } + + parsed, err := parseHiveConnectionString(params.ConnectionString) + if err != nil { + return connectionConfig{}, err + } + if !hasStructuredEndpoint { + if config.Database == "" && parsed.database != "" { + config.Database = parsed.database + } + if config.Username == "" && parsed.username != "" { + config.Username = parsed.username + } + if config.Password == "" && parsed.password != "" { + config.Password = parsed.password + } + } + if config.Database == "" { + config.Database = defaultHiveDatabase + } + + urlSections := parseHiveParameterSections(params.URLParams) + values := urlSections.session + hiveConfs := urlSections.hiveConfs + hiveVars := urlSections.hiveVars + if hasStructuredEndpoint { + deleteHiveParameters(values, "user", "username", "password", "ssl") + } else { + values = mergeHiveParameters(parsed.parameters, values) + hiveConfs = mergeHiveConfAssignments(parsed.hiveConfs, hiveConfs) + hiveVars = mergeHiveAssignments(parsed.hiveVars, hiveVars) + if value, exists := firstParameter(values, "user", "username"); exists { + config.Username = value + } + if value, exists := firstParameter(values, "password"); exists { + config.Password = value + } + } + if err := applyHiveParameters(&config, values, hiveConfs); err != nil { + return connectionConfig{}, err + } + if isZooKeeperDiscovery(config.ServiceDiscoveryMode) && len(parsed.endpoints) > 0 { + // ZooKeeper discovery needs the complete endpoint list from the JDBC URL. + config.Endpoints = parsed.endpoints + } else if host := strings.TrimSpace(params.Host); host != "" { + // DBX resolves edits and transport layers before invoking the Agent. For + // direct connections that resolved endpoint must win over the persisted URL. + port := params.Port + if port <= 0 { + port = defaultHivePort + } + for _, value := range splitEndpoints(host) { + parsedEndpoint, endpointErr := parseEndpoint(value, port) + if endpointErr != nil { + return connectionConfig{}, endpointErr + } + config.Endpoints = append(config.Endpoints, parsedEndpoint) + } + } else { + config.Endpoints = parsed.endpoints + } + applyOpenSessionVariables(&config, values, hiveConfs, hiveVars) + if err := applyDelegationToken(&config, values); err != nil { + return connectionConfig{}, err + } + applyKerberosJavaOptions(&config.Kerberos, params.AgentJavaOptions) + applyZooKeeperKerberosJavaOptions(&config.ZooKeeperKerberos, params.AgentJavaOptions) + applyKerberosEnvironment(&config.Kerberos) + if err := finalizeKerberosConfig(&config); err != nil { + return connectionConfig{}, err + } + if len(config.Endpoints) == 0 { + return connectionConfig{}, errors.New("Hive host is required") + } + + tlsConfig, err := buildTLSConfig(params, values, config.Endpoints[0].Host) + if err != nil { + return connectionConfig{}, err + } + config.TLSConfig = tlsConfig + zooKeeperTLSConfig, err := buildZooKeeperTLSConfig(values) + if err != nil { + return connectionConfig{}, err + } + config.ZooKeeperTLSConfig = zooKeeperTLSConfig + return config, nil +} + +func isZooKeeperDiscovery(mode string) bool { + return strings.EqualFold(mode, "zookeeper") || strings.EqualFold(mode, "zookeeperha") +} + +type parsedHiveConnection struct { + endpoints []endpoint + database string + username string + password string + parameters map[string]string + hiveConfs map[string]string + hiveVars map[string]string +} + +type hiveParameterSections struct { + session map[string]string + hiveConfs map[string]string + hiveVars map[string]string +} + +func newParsedHiveConnection() parsedHiveConnection { + return parsedHiveConnection{ + parameters: map[string]string{}, + hiveConfs: map[string]string{}, + hiveVars: map[string]string{}, + } +} + +func parseHiveConnectionString(raw string) (parsedHiveConnection, error) { + value := strings.TrimSpace(raw) + if value == "" { + return newParsedHiveConnection(), nil + } + if strings.HasPrefix(strings.ToLower(value), "jdbc:hive2://") { + value = value[len("jdbc:hive2://"):] + } else if strings.HasPrefix(strings.ToLower(value), "hive://") { + parsedURL, err := url.Parse(value) + if err != nil { + return parsedHiveConnection{}, fmt.Errorf("invalid Hive connection string: %w", err) + } + result := newParsedHiveConnection() + if parsedURL.User != nil { + result.username = parsedURL.User.Username() + result.password, _ = parsedURL.User.Password() + } + port := defaultHivePort + if parsedURL.Port() != "" { + parsedPort, err := strconv.Atoi(parsedURL.Port()) + if err != nil { + return parsedHiveConnection{}, fmt.Errorf("invalid Hive port: %w", err) + } + port = parsedPort + } + result.endpoints = []endpoint{{Host: parsedURL.Hostname(), Port: port}} + result.database = strings.Trim(parsedURL.Path, "/") + for key, entries := range parsedURL.Query() { + if len(entries) > 0 { + setCaseInsensitive(result.parameters, key, entries[len(entries)-1]) + } + } + return result, nil + } else { + return parsedHiveConnection{}, errors.New("Hive connection string must start with jdbc:hive2:// or hive://") + } + + result := newParsedHiveConnection() + if fragment := strings.IndexByte(value, '#'); fragment >= 0 { + result.hiveVars = parseHiveAssignments(value[fragment+1:], false) + value = value[:fragment] + } + if query := strings.IndexByte(value, '?'); query >= 0 { + result.hiveConfs = parseHiveAssignments(value[query+1:], false) + value = value[:query] + } + pathStart := strings.IndexByte(value, '/') + authority := value + pathAndParams := "" + if pathStart >= 0 { + authority = value[:pathStart] + pathAndParams = value[pathStart+1:] + } + if at := strings.LastIndex(authority, "@"); at >= 0 { + credentials := authority[:at] + authority = authority[at+1:] + if colon := strings.IndexByte(credentials, ':'); colon >= 0 { + result.username, _ = url.QueryUnescape(credentials[:colon]) + result.password, _ = url.QueryUnescape(credentials[colon+1:]) + } else { + result.username, _ = url.QueryUnescape(credentials) + } + } + for _, value := range splitEndpoints(authority) { + parsedEndpoint, err := parseEndpoint(value, defaultHivePort) + if err != nil { + return parsedHiveConnection{}, err + } + result.endpoints = append(result.endpoints, parsedEndpoint) + } + if separator := strings.IndexByte(pathAndParams, ';'); separator >= 0 { + result.database = strings.TrimSpace(pathAndParams[:separator]) + result.parameters = mergeHiveParameters(result.parameters, parseHiveParameters(pathAndParams[separator+1:])) + } else { + result.database = strings.TrimSpace(pathAndParams) + } + return result, nil +} + +func splitEndpoints(value string) []string { + parts := strings.Split(value, ",") + result := make([]string, 0, len(parts)) + for _, part := range parts { + if trimmed := strings.TrimSpace(part); trimmed != "" { + result = append(result, trimmed) + } + } + return result +} + +func parseEndpoint(value string, defaultPort int) (endpoint, error) { + value = strings.TrimSpace(value) + if value == "" { + return endpoint{}, errors.New("Hive endpoint is empty") + } + host := value + port := defaultPort + if parsedHost, parsedPort, err := net.SplitHostPort(value); err == nil { + host = parsedHost + parsed, parseErr := strconv.Atoi(parsedPort) + if parseErr != nil { + return endpoint{}, fmt.Errorf("invalid Hive endpoint %q: %w", value, parseErr) + } + port = parsed + } else if strings.Count(value, ":") == 1 { + parts := strings.SplitN(value, ":", 2) + parsed, parseErr := strconv.Atoi(parts[1]) + if parseErr != nil { + return endpoint{}, fmt.Errorf("invalid Hive endpoint %q: %w", value, parseErr) + } + host = parts[0] + port = parsed + } else if strings.HasPrefix(value, "[") && strings.HasSuffix(value, "]") { + host = strings.Trim(value, "[]") + } + if strings.TrimSpace(host) == "" || port <= 0 || port > 65535 { + return endpoint{}, fmt.Errorf("invalid Hive endpoint %q", value) + } + return endpoint{Host: host, Port: port}, nil +} + +func parseHiveParameters(raw string) map[string]string { + return parseHiveAssignments(raw, false) +} + +func parseHiveAssignments(raw string, lowercaseKeys bool) map[string]string { + result := map[string]string{} + trimmed := strings.Trim(strings.TrimSpace(raw), "?#&;") + for _, part := range strings.FieldsFunc(trimmed, func(char rune) bool { return char == ';' || char == '&' }) { + part = strings.TrimSpace(part) + if part == "" { + continue + } + key, value, found := strings.Cut(part, "=") + key = strings.TrimSpace(key) + if decoded, err := url.QueryUnescape(key); err == nil { + key = decoded + } + if lowercaseKeys { + key = strings.ToLower(key) + } + if key == "" { + continue + } + if found { + if decoded, err := url.QueryUnescape(strings.TrimSpace(value)); err == nil { + value = decoded + } + } else { + value = "" + } + if lowercaseKeys { + result[key] = value + } else { + setCaseInsensitive(result, key, value) + } + } + return result +} + +func parseHiveParameterSections(raw string) hiveParameterSections { + value := strings.TrimSpace(raw) + sections := hiveParameterSections{ + session: map[string]string{}, + hiveConfs: map[string]string{}, + hiveVars: map[string]string{}, + } + if fragment := strings.IndexByte(value, '#'); fragment >= 0 { + sections.hiveVars = parseHiveAssignments(value[fragment+1:], false) + value = value[:fragment] + } + if query := strings.IndexByte(value, '?'); query >= 0 { + sections.hiveConfs = parseHiveAssignments(value[query+1:], false) + value = value[:query] + } + sections.session = parseHiveParameters(value) + return sections +} + +func mergeHiveParameters(first, second map[string]string) map[string]string { + result := make(map[string]string, len(first)+len(second)) + for key, value := range first { + setCaseInsensitive(result, key, value) + } + for key, value := range second { + setCaseInsensitive(result, key, value) + } + return result +} + +func setCaseInsensitive(values map[string]string, key, value string) { + for existing := range values { + if strings.EqualFold(existing, key) { + delete(values, existing) + } + } + values[key] = value +} + +func deleteHiveParameters(values map[string]string, keys ...string) { + for existing := range values { + for _, key := range keys { + if strings.EqualFold(existing, key) { + delete(values, existing) + break + } + } + } +} + +func mergeHiveAssignments(first, second map[string]string) map[string]string { + result := make(map[string]string, len(first)+len(second)) + for key, value := range first { + result[key] = value + } + for key, value := range second { + result[key] = value + } + return result +} + +func mergeHiveConfAssignments(first, second map[string]string) map[string]string { + result := make(map[string]string, len(first)+len(second)) + for key, value := range first { + result[canonicalHiveConfKey(key)] = value + } + for key, value := range second { + result[canonicalHiveConfKey(key)] = value + } + return result +} + +func applyHiveParameters(config *connectionConfig, values, hiveConfs map[string]string) error { + if value := parameter(values, "auth"); value != "" { + config.Auth = strings.ToUpper(value) + config.AuthExplicit = true + } + if value := firstNonEmpty(parameter(values, "transportmode"), hiveAssignmentValue(hiveConfs, "hive.server2.transport.mode")); value != "" { + config.TransportMode = strings.ToLower(value) + config.TransportModeExplicit = true + } + if value := firstNonEmpty(parameter(values, "httppath"), hiveAssignmentValue(hiveConfs, "hive.server2.thrift.http.path")); value != "" { + config.HTTPPath = strings.TrimPrefix(value, "/") + config.HTTPPathExplicit = true + } + if value := parameter(values, "servicediscoverymode"); value != "" { + config.ServiceDiscoveryMode = strings.ToLower(value) + } + if value := parameter(values, "zookeepernamespace"); value != "" { + config.ZooKeeperNamespace = strings.Trim(value, "/") + } + if strings.EqualFold(config.ServiceDiscoveryMode, "zookeeperha") && parameter(values, "zookeepernamespace") == "" { + config.ZooKeeperNamespace = "hs2ActivePassiveHA" + } + if hasParameter(values, "ssl") { + config.TLSExplicit = true + } + config.ZooKeeperAuthScheme = parameter(values, "zookeeperauthscheme") + config.ZooKeeperAuth = parameter(values, "zookeeperauth") + config.HTTPHeaders = prefixedParameters(values, "http.header.") + config.HTTPCookies = prefixedParameters(values, "http.cookie.") + config.RequestTracking = parameterBool(values, "requesttrack") + if value, exists := firstParameter(values, "cookieauth"); exists { + config.CookieAuth = !strings.EqualFold(value, "false") + } + config.CookieName = firstNonEmpty(parameter(values, "cookiename"), defaultCookieName) + config.JWT = firstNonEmpty(parameter(values, "jwt"), os.Getenv("JWT")) + config.BrowserToken = firstNonEmpty(parameter(values, "browsertoken"), parameter(values, "token")) + config.BrowserClientID = parameter(values, "browserclientidentifier") + if value := parameter(values, "browserresponseport"); value != "" { + parsed, err := strconv.Atoi(value) + if err != nil || parsed < 0 || parsed > 65535 { + return fmt.Errorf("invalid Hive browserResponsePort %q: expected 0-65535", value) + } + config.BrowserResponsePort = parsed + } + if value := parameter(values, "browserresponsetimeout"); value != "" { + parsed, err := strconv.ParseInt(value, 10, 64) + if err != nil || parsed <= 0 { + return fmt.Errorf("invalid Hive browserResponseTimeout %q: expected positive seconds", value) + } + config.BrowserResponseTimeout = time.Duration(parsed) * time.Second + } + config.BrowserDisableSSLCheck = parameterBool(values, "browserdisablesslcheck") + if strings.EqualFold(config.Auth, "JWT") && config.JWT == "" { + return errors.New("Hive JWT authentication requires jwt or the JWT environment variable") + } + if value := parameter(values, "fetchsize"); value != "" { + parsed, err := strconv.Atoi(value) + if err != nil || parsed <= 0 { + return fmt.Errorf("invalid Hive fetchSize %q: expected a positive integer", value) + } + config.FetchSize = parsed + } + if value := parameter(values, "sockettimeout"); value != "" { + parsed, err := strconv.ParseInt(value, 10, 64) + if err != nil { + return fmt.Errorf("invalid Hive socketTimeout %q: expected seconds", value) + } + if parsed > 0 { + config.SocketTimeout = time.Duration(parsed) * time.Second + } + } + if value := parameter(values, "thrift.client.max.message.size"); value != "" { + parsed, err := strconv.ParseInt(value, 10, 32) + if err != nil { + return fmt.Errorf("invalid Hive thrift.client.max.message.size %q: expected bytes", value) + } + if parsed > 0 { + config.MaxMessageSize = int32(parsed) + } + } + if value := parameter(values, "retries"); value != "" { + parsed, err := strconv.Atoi(value) + if err == nil && parsed > 0 { + config.Retries = parsed + } + } + if value := parameter(values, "retryinterval"); value != "" { + parsed, err := strconv.ParseInt(value, 10, 64) + if err == nil && parsed >= 0 { + config.RetryInterval = time.Duration(parsed) * time.Millisecond + } + } + if value := parameter(values, "initfile"); value != "" { + statements, err := readHiveInitFile(value) + if err != nil { + return err + } + config.InitStatements = statements + } + + kerberos := &config.Kerberos + kerberos.ServerPrincipal = parameter(values, "principal") + kerberos.ServerPrincipalExplicit = kerberos.ServerPrincipal != "" + kerberos.ClientPrincipal = firstNonEmpty( + parameter(values, "kerberosprincipal"), + parameter(values, "clientprincipal"), + parameter(values, "userprincipal"), + ) + kerberos.Service = firstNonEmpty(parameter(values, "service"), serviceFromPrincipal(kerberos.ServerPrincipal), kerberos.Service) + kerberos.ServerName = parameter(values, "servername") + kerberos.Realm = firstNonEmpty(parameter(values, "realm"), realmFromPrincipal(kerberos.ClientPrincipal)) + kerberos.ConfigPath = firstNonEmpty(parameter(values, "krb5conf"), parameter(values, "kerberosconfig")) + kerberos.JAASConfigPath = parameter(values, "jaasconfig") + kerberos.KeytabPath = parameter(values, "keytab") + kerberos.CCachePath = firstNonEmpty(parameter(values, "ccache"), parameter(values, "credentialcache")) + kerberos.AuthorizationID = firstNonEmpty(parameter(values, "authorizationid"), parameter(values, "proxyuser")) + kerberos.QOP = firstNonEmpty( + parameter(values, "hive.server2.thrift.sasl.qop"), + hiveAssignmentValue(hiveConfs, "hive.server2.thrift.sasl.qop"), + parameter(values, "sasl.qop"), + parameter(values, "saslqop"), + "auth", + ) + kerberos.UseKeytab = parameterBool(values, "usekeytab") || kerberos.KeytabPath != "" + kerberos.UseTicketCache = parameterBool(values, "useticketcache") || kerberos.CCachePath != "" + kerberos.UseSSPI = parameterBool(values, "usesspi") + if hasParameter(values, "kerberosenablecanonicalhostnamecheck") { + kerberos.CanonicalHostname = parameterBool(values, "kerberosenablecanonicalhostnamecheck") + } + kerberos.ChannelBinding = parameterBool(values, "kerberoschannelbinding") || + parameterBool(values, "tlschannelbinding") || + parameterBool(values, "channelbinding") + kerberos.DisablePAFXFAST = parameterBool(values, "disablepafxfast") + if kerberos.ServerPrincipal != "" || strings.EqualFold(config.Auth, "KERBEROS") { + kerberos.Enabled = true + config.Auth = "KERBEROS" + } + + zooKeeperKerberos := &config.ZooKeeperKerberos + zooKeeperKerberos.Enabled = kerberos.ServerPrincipalExplicit + if value, exists := firstParameter(values, "hive.zookeeper.use.kerberos", "hiveconf:hive.zookeeper.use.kerberos"); exists { + zooKeeperKerberos.Enabled = booleanValue(value) + } else if value := hiveAssignmentValue(hiveConfs, "hive.zookeeper.use.kerberos"); value != "" { + zooKeeperKerberos.Enabled = booleanValue(value) + } + if value, exists := firstParameter(values, "zookeeper.sasl.client"); exists && !booleanValue(value) { + zooKeeperKerberos.Enabled = false + } + zooKeeperKerberos.Service = firstNonEmpty(parameter(values, "zookeeper.sasl.client.username"), "zookeeper") + zooKeeperKerberos.ServerPrincipal = parameter(values, "zookeeper.server.principal") + zooKeeperKerberos.Realm = parameter(values, "zookeeper.server.realm") + if hasParameter(values, "zookeeper.sasl.client.canonicalize.hostname") { + zooKeeperKerberos.CanonicalHostname = parameterBool(values, "zookeeper.sasl.client.canonicalize.hostname") + } + + return nil +} + +func applyOpenSessionVariables(config *connectionConfig, values, hiveConfs, hiveVars map[string]string) { + config.HiveConfiguration["set:hiveconf:"+resultSetUniqueColumnNames] = "false" + for key, value := range values { + lowerKey := strings.ToLower(key) + switch { + case strings.HasPrefix(lowerKey, "hiveconf:"): + config.HiveConfiguration["set:hiveconf:"+canonicalHiveConfKey(key[len("hiveconf:"):])] = value + case strings.HasPrefix(lowerKey, "hivevar:"): + config.HiveConfiguration["set:hivevar:"+key[len("hivevar:"):]] = value + } + } + for key, value := range hiveConfs { + if strings.EqualFold(key, "hive.server2.transport.mode") || strings.EqualFold(key, "hive.server2.thrift.http.path") { + continue + } + config.HiveConfiguration["set:hiveconf:"+canonicalHiveConfKey(key)] = value + } + for key, value := range hiveVars { + config.HiveConfiguration["set:hivevar:"+key] = value + } + if proxyUser := firstNonEmpty(parameter(values, "proxyuser"), parameter(values, "hive.server2.proxy.user")); proxyUser != "" { + config.HiveConfiguration["hive.server2.proxy.user"] = proxyUser + } + if value := parameter(values, "hivecreateasexternallegacy"); value != "" { + config.HiveConfiguration["set:hiveconf:hive.create.as.external.legacy"] = strings.ToLower(value) + } + if value := parameter(values, "wmpool"); value != "" { + config.HiveConfiguration["set:hivevar:wmpool"] = value + } + if value := firstNonEmpty(parameter(values, "applicationname"), parameter(values, "ApplicationName")); value != "" { + config.HiveConfiguration["set:hivevar:wmapp"] = value + } +} + +func canonicalHiveConfKey(key string) string { + if strings.EqualFold(key, resultSetUniqueColumnNames) { + return resultSetUniqueColumnNames + } + return key +} + +func hiveAssignmentValue(values map[string]string, key string) string { + for candidate, value := range values { + if strings.EqualFold(strings.TrimSpace(candidate), key) { + return value + } + } + return "" +} + +func applyDelegationToken(config *connectionConfig, values map[string]string) error { + if !strings.EqualFold(config.Auth, "DELEGATIONTOKEN") && !strings.EqualFold(config.Auth, "DELEGATION_TOKEN") { + return nil + } + token := firstNonEmpty(parameter(values, "delegationtoken"), parameter(values, "token"), config.Password) + if token == "" { + return errors.New("Hive delegation token authentication requires delegationToken, token, or password") + } + config.DelegationToken = token + identifier, password, err := decodeHadoopDelegationToken(token) + if err != nil { + return fmt.Errorf("decode Hive delegation token: %w", err) + } + config.Username = base64.StdEncoding.EncodeToString(identifier) + config.Password = base64.StdEncoding.EncodeToString(password) + return nil +} + +func decodeHadoopDelegationToken(value string) ([]byte, []byte, error) { + encoded := strings.Join(strings.Fields(strings.TrimSpace(value)), "") + if encoded == "" { + return nil, nil, errors.New("token is empty") + } + var decoded []byte + var decodeErr error + for _, encoding := range []*base64.Encoding{ + base64.RawURLEncoding, + base64.URLEncoding, + base64.RawStdEncoding, + base64.StdEncoding, + } { + decoded, decodeErr = encoding.DecodeString(encoded) + if decodeErr == nil { + break + } + } + if decodeErr != nil { + return nil, nil, decodeErr + } + reader := strings.NewReader(string(decoded)) + identifier, err := readHadoopByteArray(reader) + if err != nil { + return nil, nil, fmt.Errorf("identifier: %w", err) + } + password, err := readHadoopByteArray(reader) + if err != nil { + return nil, nil, fmt.Errorf("password: %w", err) + } + if len(identifier) == 0 || len(password) == 0 { + return nil, nil, errors.New("token identifier and password must be non-empty") + } + if _, err := readHadoopByteArray(reader); err != nil { + return nil, nil, fmt.Errorf("kind: %w", err) + } + if _, err := readHadoopByteArray(reader); err != nil { + return nil, nil, fmt.Errorf("service: %w", err) + } + if reader.Len() != 0 { + return nil, nil, errors.New("token contains trailing data") + } + return identifier, password, nil +} + +func readHadoopByteArray(reader io.ByteReader) ([]byte, error) { + length, err := readHadoopVInt(reader) + if err != nil { + return nil, err + } + if length < 0 { + return nil, fmt.Errorf("negative length %d", length) + } + if length > 64*1024*1024 { + return nil, fmt.Errorf("length %d exceeds limit", length) + } + value := make([]byte, int(length)) + byteReader, ok := reader.(io.Reader) + if !ok { + return nil, errors.New("reader cannot read token payload") + } + if _, err := io.ReadFull(byteReader, value); err != nil { + return nil, err + } + return value, nil +} + +func readHadoopVInt(reader io.ByteReader) (int64, error) { + firstByte, err := reader.ReadByte() + if err != nil { + return 0, err + } + first := int8(firstByte) + if first >= -112 { + return int64(first), nil + } + length := -111 - int(first) + negative := false + if first < -120 { + length = -119 - int(first) + negative = true + } + var value int64 + for index := 0; index < length-1; index++ { + current, readErr := reader.ReadByte() + if readErr != nil { + return 0, readErr + } + value = value<<8 | int64(current) + } + if negative { + value = ^value + } + return value, nil +} + +func applyKerberosJavaOptions(config *kerberosConfig, options []string) { + for _, option := range options { + trimmed := strings.TrimSpace(option) + switch { + case strings.HasPrefix(trimmed, "-Djava.security.krb5.conf="): + config.ConfigPath = javaSystemPropertyValue(strings.TrimPrefix(trimmed, "-Djava.security.krb5.conf=")) + case strings.HasPrefix(trimmed, "-Djava.security.auth.login.config="): + config.JAASConfigPath = javaSystemPropertyValue(strings.TrimPrefix(trimmed, "-Djava.security.auth.login.config=")) + } + } +} + +func applyZooKeeperKerberosJavaOptions(config *zooKeeperKerberosConfig, options []string) { + for _, option := range options { + trimmed := strings.TrimSpace(option) + keyValue := strings.TrimPrefix(trimmed, "-D") + key, value, found := strings.Cut(keyValue, "=") + if !strings.HasPrefix(trimmed, "-D") || !found { + continue + } + value = javaSystemPropertyValue(value) + switch strings.ToLower(strings.TrimSpace(key)) { + case "hive.zookeeper.use.kerberos": + config.Enabled = booleanValue(value) + case "zookeeper.sasl.client": + if !booleanValue(value) { + config.Enabled = false + } + case "zookeeper.sasl.client.username": + config.Service = firstNonEmpty(value, "zookeeper") + case "zookeeper.sasl.client.canonicalize.hostname": + config.CanonicalHostname = booleanValue(value) + case "zookeeper.server.principal": + config.ServerPrincipal = strings.TrimSpace(value) + case "zookeeper.server.realm": + config.Realm = strings.TrimSpace(value) + } + } +} + +func javaSystemPropertyValue(value string) string { + value = strings.TrimSpace(value) + if len(value) >= 2 && value[0] == '"' && value[len(value)-1] == '"' { + return value[1 : len(value)-1] + } + return value +} + +func applyKerberosEnvironment(config *kerberosConfig) { + config.ConfigPath = firstNonEmpty(config.ConfigPath, os.Getenv("KRB5_CONFIG")) + config.CCachePath = firstNonEmpty(config.CCachePath, os.Getenv("KRB5CCNAME")) + config.KeytabPath = firstNonEmpty(config.KeytabPath, os.Getenv("KRB5_CLIENT_KTNAME"), os.Getenv("KRB5_KTNAME")) +} + +func finalizeKerberosConfig(config *connectionConfig) error { + kerberos := &config.Kerberos + if !kerberos.Enabled { + return nil + } + kerberos.Password = config.Password + kerberos.ConfigPath = normalizeKerberosReference(kerberos.ConfigPath) + kerberos.CCachePath = normalizeKerberosReference(kerberos.CCachePath) + kerberos.KeytabPath = normalizeKerberosReference(kerberos.KeytabPath) + kerberos.JAASConfigPath = normalizeKerberosReference(kerberos.JAASConfigPath) + if kerberos.JAASConfigPath != "" { + if err := applyKerberosJAASFile(kerberos); err != nil { + return err + } + kerberos.KeytabPath = normalizeKerberosReference(kerberos.KeytabPath) + kerberos.CCachePath = normalizeKerberosReference(kerberos.CCachePath) + } + if kerberos.ConfigPath == "" { + if candidate := defaultKerberosConfigPath(); fileExists(candidate) { + kerberos.ConfigPath = candidate + } + } + if !kerberos.UseTicketCache && kerberos.CCachePath == "" { + if candidate := defaultKerberosCCachePath(); fileExists(candidate) { + kerberos.CCachePath = candidate + kerberos.UseTicketCache = true + } + } + if runtime.GOOS == "windows" && kerberos.ConfigPath == "" && kerberos.KeytabPath == "" && kerberos.CCachePath == "" { + kerberos.UseSSPI = true + } + if kerberos.UseSSPI { + return nil + } + if kerberos.ConfigPath == "" { + return errors.New("Kerberos requires krb5.conf or Windows SSPI") + } + if kerberos.ClientPrincipal == "" && !kerberos.UseTicketCache && !kerberos.UseKeytab { + kerberos.ClientPrincipal = strings.TrimSpace(config.Username) + } + if kerberos.KeytabPath != "" { + kerberos.UseKeytab = true + } + if kerberos.CCachePath != "" { + kerberos.UseTicketCache = true + } + kerberos.Realm = firstNonEmpty(kerberos.Realm, realmFromPrincipal(kerberos.ClientPrincipal)) + if !kerberos.UseTicketCache && !kerberos.UseKeytab && (kerberos.ClientPrincipal == "" || kerberos.Password == "") { + return errors.New("Kerberos requires SSPI, credential cache, keytab, or principal and password") + } + return nil +} + +var jaasOptionPattern = regexp.MustCompile(`(?i)\b(principal|keytab|ticketcache|usekeytab|useticketcache)\s*=\s*("(?:\\.|[^"])*"|'(?:\\.|[^'])*'|[^\s;]+)`) + +func applyKerberosJAASFile(config *kerberosConfig) error { + contents, err := os.ReadFile(config.JAASConfigPath) + if err != nil { + return fmt.Errorf("read Kerberos JAAS config: %w", err) + } + text := string(contents) + module := strings.Index(strings.ToLower(text), "krb5loginmodule") + if module < 0 { + return errors.New("Kerberos JAAS config contains no Krb5LoginModule") + } + block := text[module:] + if end := strings.IndexByte(block, ';'); end >= 0 { + block = block[:end] + } + for _, match := range jaasOptionPattern.FindAllStringSubmatch(block, -1) { + key := strings.ToLower(match[1]) + value := decodeJAASValue(match[2]) + switch key { + case "principal": + if config.ClientPrincipal == "" { + config.ClientPrincipal = value + } + case "keytab": + if config.KeytabPath == "" { + config.KeytabPath = value + } + case "ticketcache": + if config.CCachePath == "" { + config.CCachePath = value + } + case "usekeytab": + config.UseKeytab = config.UseKeytab || parseJAASBool(value) + case "useticketcache": + config.UseTicketCache = config.UseTicketCache || parseJAASBool(value) + } + } + return nil +} + +func decodeJAASValue(value string) string { + value = strings.TrimSpace(value) + if len(value) >= 2 && ((value[0] == '"' && value[len(value)-1] == '"') || (value[0] == '\'' && value[len(value)-1] == '\'')) { + value = value[1 : len(value)-1] + } + value = strings.ReplaceAll(value, `\\`, `\`) + value = strings.ReplaceAll(value, `\"`, `"`) + value = strings.ReplaceAll(value, `\'`, `'`) + return value +} + +func parseJAASBool(value string) bool { + switch strings.ToLower(strings.TrimSpace(value)) { + case "1", "true", "yes", "on": + return true + default: + return false + } +} + +func fileExists(path string) bool { + if strings.TrimSpace(path) == "" { + return false + } + info, err := os.Stat(path) + return err == nil && !info.IsDir() +} + +func normalizeKerberosReference(value string) string { + value = strings.TrimSpace(value) + if value == "" { + return "" + } + if strings.HasPrefix(strings.ToUpper(value), "FILE:") { + value = value[5:] + } + if strings.HasPrefix(value, "~/") { + if home, err := os.UserHomeDir(); err == nil { + value = filepath.Join(home, value[2:]) + } + } + return filepath.Clean(value) +} + +func buildTLSConfig(params connectParams, values map[string]string, serverName string) (*tls.Config, error) { + enabled := params.SSL || parameterBool(values, "ssl") || strings.EqualFold(parameter(values, "ssl"), "true") + if !enabled { + return nil, nil + } + config := &tls.Config{MinVersion: tls.VersionTLS12, ServerName: serverName} + if parameterBool(values, "sslinsecureskipverify") || parameterBool(values, "allowselfsigned") { + config.InsecureSkipVerify = true + } + var customRoots *x509.CertPool + credentialProviderPath := parameter(values, "storepasswordpath") + if path := strings.TrimSpace(params.CACertPath); path != "" { + contents, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("read Hive CA certificate: %w", err) + } + customRoots = x509.NewCertPool() + if !customRoots.AppendCertsFromPEM(contents) { + return nil, errors.New("Hive CA certificate contains no certificates") + } + } + trustStoreLocation := parameter(values, "ssltruststore") + if trustStoreLocation != "" { + if parameter(values, "truststorepassword") == "" && credentialProviderPath != "" { + return nil, errors.New("Hive storePasswordPath uses the Java Hadoop credential-provider format; configure trustStorePassword explicitly for the native agent") + } + certificates, err := loadTrustStore( + trustStoreLocation, + parameter(values, "truststorepassword"), + parameter(values, "truststoretype"), + ) + if err != nil { + return nil, fmt.Errorf("load Hive truststore: %w", err) + } + if customRoots == nil { + customRoots = x509.NewCertPool() + } + for _, certificate := range certificates { + customRoots.AddCert(certificate) + } + } + config.RootCAs = customRoots + if params.ClientCertPath != "" || params.ClientKeyPath != "" { + if params.ClientCertPath == "" || params.ClientKeyPath == "" { + return nil, errors.New("Hive client certificate and key must be configured together") + } + certificate, err := tls.LoadX509KeyPair(params.ClientCertPath, params.ClientKeyPath) + if err != nil { + return nil, fmt.Errorf("load Hive client certificate: %w", err) + } + config.Certificates = []tls.Certificate{certificate} + } + keyStoreLocation := parameter(values, "sslkeystore") + if keyStoreLocation != "" { + if parameter(values, "keystorepassword") == "" && credentialProviderPath != "" { + return nil, errors.New("Hive storePasswordPath uses the Java Hadoop credential-provider format; configure keyStorePassword explicitly for the native agent") + } + certificate, err := loadClientKeyStore( + keyStoreLocation, + parameter(values, "keystorepassword"), + parameter(values, "keystoretype"), + ) + if err != nil { + return nil, fmt.Errorf("load Hive keystore: %w", err) + } + config.Certificates = append(config.Certificates, certificate) + } + if parameterBool(values, "twoway") { + if keyStoreLocation == "" && len(config.Certificates) == 0 { + return nil, errors.New("Hive two-way TLS requires sslKeyStore or a client certificate") + } + if trustStoreLocation == "" && config.RootCAs == nil { + return nil, errors.New("Hive two-way TLS requires sslTrustStore or a CA certificate") + } + } + return config, nil +} + +func parameter(values map[string]string, key string) string { + for candidate, value := range values { + if strings.EqualFold(strings.TrimSpace(candidate), key) { + return strings.TrimSpace(value) + } + } + return "" +} + +func parameterBool(values map[string]string, key string) bool { + return booleanValue(parameter(values, key)) +} + +func booleanValue(value string) bool { + value = strings.ToLower(strings.TrimSpace(value)) + return value == "1" || value == "true" || value == "yes" || value == "on" +} + +func firstParameter(values map[string]string, keys ...string) (string, bool) { + for _, key := range keys { + for candidate, value := range values { + if strings.EqualFold(strings.TrimSpace(candidate), key) { + return strings.TrimSpace(value), true + } + } + } + return "", false +} + +func hasParameter(values map[string]string, key string) bool { + _, exists := firstParameter(values, key) + return exists +} + +func prefixedParameters(values map[string]string, prefix string) map[string]string { + result := map[string]string{} + for key, value := range values { + if len(key) <= len(prefix) || !strings.EqualFold(key[:len(prefix)], prefix) { + continue + } + name := strings.TrimSpace(key[len(prefix):]) + if name != "" { + result[name] = value + } + } + return result +} + +func readHiveInitFile(path string) ([]string, error) { + file, err := os.Open(path) + if err != nil { + return nil, fmt.Errorf("read Hive initFile: %w", err) + } + defer file.Close() + + var script strings.Builder + scanner := bufio.NewScanner(file) + scanner.Buffer(make([]byte, 0, 64*1024), 16*1024*1024) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, "--") { + continue + } + script.WriteString(line) + script.WriteByte(' ') + } + if err := scanner.Err(); err != nil { + return nil, fmt.Errorf("read Hive initFile: %w", err) + } + + statements := make([]string, 0) + for _, statement := range strings.Split(script.String(), ";") { + if trimmed := strings.TrimSpace(statement); trimmed != "" { + statements = append(statements, trimmed) + } + } + return statements, nil +} + +func firstNonEmpty(values ...string) string { + for _, value := range values { + if trimmed := strings.TrimSpace(value); trimmed != "" { + return trimmed + } + } + return "" +} + +func serviceFromPrincipal(principal string) string { + value := strings.TrimSpace(principal) + if separator := strings.IndexByte(value, '/'); separator > 0 { + return value[:separator] + } + return "" +} + +func hostFromPrincipal(principal string) string { + value := strings.TrimSpace(principal) + separator := strings.IndexByte(value, '/') + if separator < 0 { + return "" + } + value = value[separator+1:] + if realm := strings.IndexByte(value, '@'); realm >= 0 { + value = value[:realm] + } + if value == "_HOST" { + return "" + } + return value +} + +func realmFromPrincipal(principal string) string { + if separator := strings.LastIndexByte(principal, '@'); separator >= 0 { + return strings.TrimSpace(principal[separator+1:]) + } + return "" +} diff --git a/agents/drivers/argo-go/config_test.go b/agents/drivers/argo-go/config_test.go new file mode 100644 index 0000000000..1e602a5853 --- /dev/null +++ b/agents/drivers/argo-go/config_test.go @@ -0,0 +1,713 @@ +package main + +import ( + "bytes" + "crypto/x509" + "encoding/base64" + "net/url" + "os" + "path/filepath" + "reflect" + "runtime" + "strings" + "testing" + "time" + + pkcs12 "software.sslmate.com/src/go-pkcs12" +) + +func writeHadoopVInt(buffer *bytes.Buffer, value int64) { + if value >= -112 && value <= 127 { + buffer.WriteByte(byte(int8(value))) + return + } + lengthMarker := int8(-112) + encoded := value + if value < 0 { + encoded = ^value + lengthMarker = -120 + } + temporary := encoded + for temporary != 0 { + temporary >>= 8 + lengthMarker-- + } + buffer.WriteByte(byte(lengthMarker)) + length := -int(lengthMarker) + if lengthMarker < -120 { + length -= 120 + } else { + length -= 112 + } + for index := length; index != 0; index-- { + shift := uint((index - 1) * 8) + buffer.WriteByte(byte(encoded >> shift)) + } +} + +func encodeHadoopToken(identifier, password, kind, service []byte) string { + var buffer bytes.Buffer + for _, value := range [][]byte{identifier, password, kind, service} { + writeHadoopVInt(&buffer, int64(len(value))) + buffer.Write(value) + } + return base64.RawURLEncoding.EncodeToString(buffer.Bytes()) +} + +func TestParseDirectJDBCConnection(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hive.example.com:10001/analytics;transportMode=http;httpPath=gateway;ssl=true", + Username: "alice", + Password: "secret", + }) + if err != nil { + t.Fatal(err) + } + if len(config.Endpoints) != 1 || config.Endpoints[0] != (endpoint{Host: "hive.example.com", Port: 10001}) { + t.Fatalf("unexpected endpoints: %#v", config.Endpoints) + } + if config.Database != "analytics" || config.TransportMode != "http" || config.HTTPPath != "gateway" { + t.Fatalf("unexpected config: %#v", config) + } + if config.TLSConfig == nil { + t.Fatal("expected TLS config") + } +} + +func TestParseHTTPKerberosChannelBinding(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hive.example.com:10001/default;transportMode=http;ssl=true;auth=KERBEROS;principal=HTTP/_HOST@EXAMPLE.COM;kerberosChannelBinding=true", + Username: "alice@EXAMPLE.COM", + Password: "secret", + AgentJavaOptions: []string{"-Djava.security.krb5.conf=/etc/krb5.conf"}, + }) + if err != nil { + t.Fatal(err) + } + if !config.Kerberos.ChannelBinding { + t.Fatal("expected HTTP Kerberos TLS channel binding") + } +} + +func TestParseZooKeeperKerberosJDBCConnection(t *testing.T) { + configPath := "/etc/krb5.conf" + if runtime.GOOS == "windows" { + configPath = `C:\ProgramData\MIT\Kerberos5\krb5.ini` + } + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://zk1.example.com:2181,zk2.example.com:2181/default;serviceDiscoveryMode=zooKeeper;zooKeeperNamespace=kyuubi;principal=hive/_HOST@EXAMPLE.COM;hive.server2.thrift.sasl.qop=auth-conf", + Username: "alice@EXAMPLE.COM", + Password: "secret", + AgentJavaOptions: []string{"-Djava.security.krb5.conf=" + configPath}, + }) + if err != nil { + t.Fatal(err) + } + if config.ServiceDiscoveryMode != "zookeeper" || config.ZooKeeperNamespace != "kyuubi" { + t.Fatalf("unexpected ZooKeeper config: %#v", config) + } + if len(config.Endpoints) != 2 || config.Endpoints[1].Host != "zk2.example.com" || config.Endpoints[1].Port != 2181 { + t.Fatalf("unexpected endpoints: %#v", config.Endpoints) + } + if !config.Kerberos.Enabled || config.Kerberos.Service != "hive" || config.Kerberos.QOP != "auth-conf" { + t.Fatalf("unexpected Kerberos config: %#v", config.Kerberos) + } + if config.Kerberos.ConfigPath != configPath { + t.Fatalf("unexpected krb5 config path: %q", config.Kerberos.ConfigPath) + } + if !config.ZooKeeperKerberos.Enabled || config.ZooKeeperKerberos.Service != "zookeeper" || !config.ZooKeeperKerberos.CanonicalHostname { + t.Fatalf("unexpected ZooKeeper Kerberos config: %#v", config.ZooKeeperKerberos) + } +} + +func TestParseZooKeeperKerberosCompatibilityProperties(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://zk.example.com:2181/default;serviceDiscoveryMode=zooKeeper;principal=hive/_HOST@HIVE.EXAMPLE.COM;zookeeper.sasl.client.username=zkservice;zookeeper.server.realm=ZK.EXAMPLE.COM;zookeeper.sasl.client.canonicalize.hostname=false", + Username: "alice@EXAMPLE.COM", + Password: "secret", + AgentJavaOptions: []string{ + "-Djava.security.krb5.conf=/etc/krb5.conf", + "-Dzookeeper.server.principal=zookeeper/zk.example.com@EXPLICIT.EXAMPLE.COM", + }, + }) + if err != nil { + t.Fatal(err) + } + want := zooKeeperKerberosConfig{ + Enabled: true, + Service: "zkservice", + ServerPrincipal: "zookeeper/zk.example.com@EXPLICIT.EXAMPLE.COM", + Realm: "ZK.EXAMPLE.COM", + CanonicalHostname: false, + } + if config.ZooKeeperKerberos != want { + t.Fatalf("ZooKeeper Kerberos config = %#v, want %#v", config.ZooKeeperKerberos, want) + } +} + +func TestZooKeeperKerberosCanBeDisabledExplicitly(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://zk.example.com:2181/default;serviceDiscoveryMode=zooKeeper;principal=hive/_HOST@EXAMPLE.COM;hive.zookeeper.use.kerberos=true", + Username: "alice@EXAMPLE.COM", + Password: "secret", + AgentJavaOptions: []string{ + "-Djava.security.krb5.conf=/etc/krb5.conf", + "-Dzookeeper.sasl.client=false", + }, + }) + if err != nil { + t.Fatal(err) + } + if config.ZooKeeperKerberos.Enabled { + t.Fatalf("ZooKeeper Kerberos should be disabled: %#v", config.ZooKeeperKerberos) + } +} + +func TestURLParamsOverrideConnectionString(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + Host: "hive.example.com", + Port: 10000, + ConnectionString: "jdbc:hive2://old.example.com:10000/default;transportMode=binary", + URLParams: "transportMode=http;httpPath=proxy", + }) + if err != nil { + t.Fatal(err) + } + if config.TransportMode != "http" || config.HTTPPath != "proxy" { + t.Fatalf("URL params did not override connection string: %#v", config) + } + if len(config.Endpoints) != 1 || config.Endpoints[0] != (endpoint{Host: "hive.example.com", Port: 10000}) { + t.Fatalf("resolved form endpoint did not override connection string: %#v", config.Endpoints) + } +} + +func TestResolvedDirectEndpointAndDatabasePreserveJDBCParameterSections(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + Host: "127.0.0.1", + Port: 18080, + Database: "analytics", + ConnectionString: "jdbc:hive2://old.example.com:10000/default;transportMode=http;httpPath=gateway?hive.exec.dynamic.partition=true#SourceTable=events", + URLParams: "transportMode=http;httpPath=gateway?hive.exec.dynamic.partition=true#SourceTable=events", + }) + if err != nil { + t.Fatal(err) + } + if len(config.Endpoints) != 1 || config.Endpoints[0] != (endpoint{Host: "127.0.0.1", Port: 18080}) { + t.Fatalf("unexpected resolved endpoint: %#v", config.Endpoints) + } + if config.Database != "analytics" { + t.Fatalf("resolved database = %q, want analytics", config.Database) + } + if config.TransportMode != "http" || config.HTTPPath != "gateway" { + t.Fatalf("JDBC session parameters were not preserved: %#v", config) + } + if config.HiveConfiguration["set:hiveconf:hive.exec.dynamic.partition"] != "true" { + t.Fatalf("JDBC hiveConfs were not preserved: %#v", config.HiveConfiguration) + } + if config.HiveConfiguration["set:hivevar:SourceTable"] != "events" { + t.Fatalf("JDBC hiveVars were not preserved: %#v", config.HiveConfiguration) + } +} + +func TestStructuredFieldsOverridePersistedJDBCValues(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + Host: "127.0.0.1", + Port: 18080, + Username: "edited-user", + Password: "edited-password", + ConnectionString: "jdbc:hive2://old-user:old-password@old.example.com:10000/old_database;transportMode=http;auth=LDAP;ssl=true;serviceDiscoveryMode=zooKeeper?hive.exec.dynamic.partition=true#SourceTable=events", + URLParams: "user=url-user;password=url-password;ssl=true?hive.exec.dynamic.partition=true#SourceTable=events", + }) + if err != nil { + t.Fatal(err) + } + if len(config.Endpoints) != 1 || config.Endpoints[0] != (endpoint{Host: "127.0.0.1", Port: 18080}) { + t.Fatalf("unexpected resolved endpoint: %#v", config.Endpoints) + } + if config.Database != defaultHiveDatabase { + t.Fatalf("database = %q, want %q", config.Database, defaultHiveDatabase) + } + if config.Username != "edited-user" || config.Password != "edited-password" { + t.Fatalf("structured credentials were overwritten: username=%q password=%q", config.Username, config.Password) + } + if config.TransportMode != "binary" || config.Auth != "NONE" || config.ServiceDiscoveryMode != "" { + t.Fatalf("removed JDBC parameters remained active: %#v", config) + } + if config.TLSConfig != nil { + t.Fatal("disabled structured SSL was re-enabled by persisted JDBC parameters") + } + if config.HiveConfiguration["set:hiveconf:hive.exec.dynamic.partition"] != "true" { + t.Fatalf("current hiveConfs were not preserved: %#v", config.HiveConfiguration) + } + if config.HiveConfiguration["set:hivevar:SourceTable"] != "events" { + t.Fatalf("current hiveVars were not preserved: %#v", config.HiveConfiguration) + } +} + +func TestConnectionStringOnlyKeepsJDBCValues(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://raw-user:raw-password@hive.example.com:10001/analytics;transportMode=http;ssl=true", + }) + if err != nil { + t.Fatal(err) + } + if config.Database != "analytics" || config.Username != "raw-user" || config.Password != "raw-password" { + t.Fatalf("connection-string-only fields were not preserved: %#v", config) + } + if config.TransportMode != "http" || config.TLSConfig == nil { + t.Fatalf("connection-string-only parameters were not preserved: %#v", config) + } +} + +func TestZooKeeperDiscoveryKeepsAllJDBCEndpoints(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + Host: "127.0.0.1", + Port: 12181, + Database: "analytics", + ConnectionString: "jdbc:hive2://zk1.example.com:2181,zk2.example.com:2181/default;serviceDiscoveryMode=zooKeeper;zooKeeperNamespace=hiveserver2", + URLParams: "serviceDiscoveryMode=zooKeeper;zooKeeperNamespace=hiveserver2", + }) + if err != nil { + t.Fatal(err) + } + if len(config.Endpoints) != 2 || config.Endpoints[0].Host != "zk1.example.com" || config.Endpoints[1].Host != "zk2.example.com" { + t.Fatalf("ZooKeeper discovery endpoints were not preserved: %#v", config.Endpoints) + } + if config.Database != "analytics" { + t.Fatalf("resolved database = %q, want analytics", config.Database) + } +} + +func TestImpalaDefaultsToNoSASL(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + DatabaseType: "impala", + Host: "impala.example.com", + Port: 21050, + Database: "analytics", + }) + if err != nil { + t.Fatal(err) + } + if config.Auth != "NOSASL" { + t.Fatalf("unexpected Impala auth mode: %q", config.Auth) + } + if config.Kerberos.Service != "impala" { + t.Fatalf("unexpected Impala Kerberos service: %q", config.Kerberos.Service) + } +} + +func TestImpalaExplicitAuthenticationOverridesDefaults(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + DatabaseType: "impala", + Host: "impala.example.com", + Port: 21050, + URLParams: "auth=NONE;service=custom", + }) + if err != nil { + t.Fatal(err) + } + if config.Auth != "NONE" || config.Kerberos.Service != "custom" { + t.Fatalf("explicit Impala authentication was not preserved: %#v", config.Kerberos) + } +} + +func TestImpalaLDAPHTTPSSLConfiguration(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + DatabaseType: "impala", + Host: "impala.example.com", + Port: 21050, + Username: "alice", + Password: "secret", + URLParams: "auth=LDAP;transportMode=http;httpPath=cliservice", + SSL: true, + }) + if err != nil { + t.Fatal(err) + } + if config.Auth != "LDAP" || config.TransportMode != "http" || config.HTTPPath != "cliservice" { + t.Fatalf("unexpected Impala LDAP transport config: %#v", config) + } + if config.Username != "alice" || config.Password != "secret" { + t.Fatalf("Impala LDAP credentials were not preserved: %q / %q", config.Username, config.Password) + } + if config.TLSConfig == nil || config.TLSConfig.ServerName != "impala.example.com" { + t.Fatalf("unexpected Impala LDAP TLS config: %#v", config.TLSConfig) + } + if config.Kerberos.Service != defaultImpalaService { + t.Fatalf("unexpected Impala service default: %q", config.Kerberos.Service) + } +} + +func TestParseStandardJDBCURLSectionsAndCredentials(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com:10001/analytics;user=alice;password=p%40ss?hive.server2.transport.mode=http;hive.server2.thrift.http.path=proxy;hive.exec.dynamic.partition=true#SourceTable=events", + URLParams: "user=bob?hive.exec.dynamic.partition=false#SourceTable=override", + }) + if err != nil { + t.Fatal(err) + } + if config.Username != "bob" || config.Password != "p@ss" { + t.Fatalf("unexpected credentials: %q / %q", config.Username, config.Password) + } + if config.TransportMode != "http" || config.HTTPPath != "proxy" { + t.Fatalf("deprecated Hive conf transport settings were not applied: %#v", config) + } + want := map[string]string{ + "set:hiveconf:hive.resultset.use.unique.column.names": "false", + "set:hiveconf:hive.exec.dynamic.partition": "false", + "set:hivevar:SourceTable": "override", + } + if !reflect.DeepEqual(config.HiveConfiguration, want) { + t.Fatalf("unexpected OpenSession configuration: %#v", config.HiveConfiguration) + } +} + +func TestOpenSessionCompatibilityVariablesFromSessionParams(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + Host: "hs2.example.com", + URLParams: "proxyUser=alice;hiveCreateAsExternalLegacy=TRUE;wmPool=etl;hiveconf:hive.exec.compress.output=true;hivevar:source=events", + }) + if err != nil { + t.Fatal(err) + } + want := map[string]string{ + "hive.server2.proxy.user": "alice", + "set:hiveconf:hive.create.as.external.legacy": "true", + "set:hiveconf:hive.exec.compress.output": "true", + "set:hiveconf:hive.resultset.use.unique.column.names": "false", + "set:hivevar:source": "events", + "set:hivevar:wmpool": "etl", + } + if !reflect.DeepEqual(config.HiveConfiguration, want) { + t.Fatalf("unexpected OpenSession compatibility variables: %#v", config.HiveConfiguration) + } +} + +func TestOpenSessionUsesLeafResultLabelsUnlessExplicitlyOverridden(t *testing.T) { + tests := []struct { + name string + connection string + urlParams string + want string + }{ + {name: "default", want: "false"}, + { + name: "JDBC hiveconf override", + connection: "jdbc:hive2://hs2.example.com:10000/default?hive.resultset.use.unique.column.names=true", + want: "true", + }, + { + name: "URL hiveconf override", + urlParams: "?HIVE.RESULTSET.USE.UNIQUE.COLUMN.NAMES=true", + want: "true", + }, + { + name: "URL hiveconf wins over JDBC hiveconf", + connection: "jdbc:hive2://hs2.example.com:10000/default?hive.resultset.use.unique.column.names=true", + urlParams: "?hive.resultset.use.unique.column.names=false", + want: "false", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + params := connectParams{Host: "hs2.example.com", ConnectionString: test.connection, URLParams: test.urlParams} + if test.connection != "" { + params.Host = "" + } + config, err := parseConnectionConfig(params) + if err != nil { + t.Fatal(err) + } + const key = "set:hiveconf:hive.resultset.use.unique.column.names" + if got := config.HiveConfiguration[key]; got != test.want { + t.Fatalf("%s = %q, want %q; config=%#v", key, got, test.want, config.HiveConfiguration) + } + matches := 0 + for candidate := range config.HiveConfiguration { + if strings.EqualFold(candidate, key) { + matches++ + } + } + if matches != 1 { + t.Fatalf("result label hiveconf must appear exactly once: %#v", config.HiveConfiguration) + } + }) + } +} + +func TestHiveAssignmentMergeCanonicalizesOnlyResultLabelSetting(t *testing.T) { + first := map[string]string{ + "CaseSensitive": "first", + "HIVE.RESULTSET.USE.UNIQUE.COLUMN.NAMES": "true", + } + second := map[string]string{ + "casesensitive": "second", + "hive.resultset.use.unique.column.names": "false", + } + + hiveConfs := mergeHiveConfAssignments(first, second) + if got := hiveConfs[resultSetUniqueColumnNames]; got != "false" { + t.Fatalf("result-label setting override = %q, want false; values=%#v", got, hiveConfs) + } + if got := len(hiveConfs); got != 3 { + t.Fatalf("unrelated case-distinct Hive confs were collapsed: %#v", hiveConfs) + } + + hiveVars := mergeHiveAssignments(first, second) + if got := len(hiveVars); got != 4 { + t.Fatalf("case-distinct Hive variables were collapsed: %#v", hiveVars) + } +} + +func TestParseHiveJDBCClientCompatibilityOptions(t *testing.T) { + directory := t.TempDir() + initPath := filepath.Join(directory, "hive-init.sql") + if err := os.WriteFile(initPath, []byte("# ignored\nSET hive.exec.dynamic.partition=true;\n-- ignored\nUSE analytics;\n"), 0o600); err != nil { + t.Fatal(err) + } + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com:10001/default;transportMode=http;auth=jwt;jwt=signed-token;fetchSize=77;socketTimeout=9;thrift.client.max.message.size=1048576;retries=3;retryInterval=250;requestTrack=true;cookieAuth=false;cookieName=CustomAuth;http.header.X-Trace-ID=trace-value;http.cookie.SessionID=cookie-value;applicationName=dbx-hive;initFile=" + url.QueryEscape(initPath), + }) + if err != nil { + t.Fatal(err) + } + if config.Auth != "JWT" || config.JWT != "signed-token" { + t.Fatalf("unexpected JWT config: %#v", config) + } + if config.FetchSize != 77 || config.SocketTimeout != 9*time.Second || config.MaxMessageSize != 1048576 { + t.Fatalf("unexpected client sizing config: %#v", config) + } + if config.Retries != 3 || config.RetryInterval != 250*time.Millisecond { + t.Fatalf("unexpected retry config: %#v", config) + } + if config.CookieAuth || config.CookieName != "CustomAuth" { + t.Fatalf("unexpected cookie auth config: %#v", config) + } + if config.HTTPHeaders["X-Trace-ID"] != "trace-value" || config.HTTPCookies["SessionID"] != "cookie-value" { + t.Fatalf("HTTP header or cookie case was not preserved: %#v / %#v", config.HTTPHeaders, config.HTTPCookies) + } + if !config.RequestTracking { + t.Fatal("requestTrack was not enabled") + } + if config.HiveConfiguration["set:hivevar:wmapp"] != "dbx-hive" { + t.Fatalf("applicationName was not mapped: %#v", config.HiveConfiguration) + } + if !reflect.DeepEqual(config.InitStatements, []string{"SET hive.exec.dynamic.partition=true", "USE analytics"}) { + t.Fatalf("unexpected init statements: %#v", config.InitStatements) + } +} + +func TestJWTCanComeFromEnvironment(t *testing.T) { + t.Setenv("JWT", "environment-token") + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com:10001/default;transportMode=http;auth=jwt", + }) + if err != nil { + t.Fatal(err) + } + if config.JWT != "environment-token" { + t.Fatalf("unexpected JWT: %q", config.JWT) + } +} + +func TestBrowserAuthSupportsInteractiveSSOParameters(t *testing.T) { + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com:10001/default;transportMode=http;auth=browser;browserResponsePort=18080;browserResponseTimeout=45;browserDisableSslCheck=true", + }) + if err != nil { + t.Fatal(err) + } + if config.BrowserToken != "" || config.BrowserResponsePort != 18080 || config.BrowserResponseTimeout != 45*time.Second || !config.BrowserDisableSSLCheck { + t.Fatalf("unexpected browser SSO config: %#v", config) + } +} + +func TestInvalidHiveJDBCClientSizesAreRejected(t *testing.T) { + for _, params := range []string{ + "fetchSize=zero", + "socketTimeout=soon", + "thrift.client.max.message.size=huge", + } { + _, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com/default;" + params, + }) + if err == nil { + t.Fatalf("expected %q to be rejected", params) + } + } +} + +func TestJavaCredentialProviderPasswordPathIsNotSilentlyIgnored(t *testing.T) { + _, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com/default;ssl=true;sslTrustStore=/tmp/truststore.jks;storePasswordPath=jceks://file/tmp/hive.jceks", + }) + if err == nil || !strings.Contains(err.Error(), "configure trustStorePassword explicitly") { + t.Fatalf("unexpected credential-provider error: %v", err) + } +} + +func TestKerberosAcceptsSaslQopCompatibilityAlias(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "krb5.conf") + if err := os.WriteFile(configPath, []byte("[libdefaults]\n default_realm = EXAMPLE.COM\n"), 0o600); err != nil { + t.Fatal(err) + } + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com/default;auth=kerberos;principal=hive/_HOST@EXAMPLE.COM;kerberosPrincipal=alice@EXAMPLE.COM;sasl.qop=auth-int", + Password: "secret", + AgentJavaOptions: []string{"-Djava.security.krb5.conf=" + configPath}, + }) + if err != nil { + t.Fatal(err) + } + if config.Kerberos.QOP != "auth-int" { + t.Fatalf("unexpected SASL QOP: %q", config.Kerberos.QOP) + } +} + +func TestDelegationTokenProducesDigestCredentials(t *testing.T) { + identifier := []byte("token-identifier") + password := []byte("token-password") + token := encodeHadoopToken(identifier, password, []byte("HIVE_DELEGATION_TOKEN"), []byte("hs2.example.com:10000")) + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com/default;auth=delegationToken;delegationToken=" + token, + }) + if err != nil { + t.Fatal(err) + } + if config.Auth != "DELEGATIONTOKEN" { + t.Fatalf("unexpected auth: %q", config.Auth) + } + if config.DelegationToken != token { + t.Fatalf("raw delegation token was not retained: %q", config.DelegationToken) + } + if config.Username != base64.StdEncoding.EncodeToString(identifier) || config.Password != base64.StdEncoding.EncodeToString(password) { + t.Fatalf("unexpected delegation token credentials: %q / %q", config.Username, config.Password) + } +} + +func TestDelegationTokenRejectsMalformedValue(t *testing.T) { + _, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com/default;auth=delegationToken;delegationToken=not-a-token", + }) + if err == nil || !strings.Contains(err.Error(), "decode Hive delegation token") { + t.Fatalf("expected delegation token error, got %v", err) + } +} + +func TestParseIPv6Endpoint(t *testing.T) { + value, err := parseEndpoint("[2001:db8::1]:10000", defaultHivePort) + if err != nil { + t.Fatal(err) + } + if value.Host != "2001:db8::1" || value.Port != 10000 { + t.Fatalf("unexpected endpoint: %#v", value) + } +} + +func TestKerberosSeparatesServerAndClientPrincipals(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "krb5.conf") + if err := os.WriteFile(configPath, []byte("[libdefaults]\n default_realm = CLIENT.EXAMPLE.COM\n"), 0o600); err != nil { + t.Fatal(err) + } + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://alias.example.com:10000/default;auth=kerberos;principal=hive/_HOST@SERVER.EXAMPLE.COM;kerberosPrincipal=alice@CLIENT.EXAMPLE.COM;kerberosEnableCanonicalHostnameCheck=false", + Password: "secret", + AgentJavaOptions: []string{`-Djava.security.krb5.conf="` + configPath + `"`}, + }) + if err != nil { + t.Fatal(err) + } + if config.Kerberos.ServerPrincipal != "hive/_HOST@SERVER.EXAMPLE.COM" { + t.Fatalf("unexpected server principal: %q", config.Kerberos.ServerPrincipal) + } + if config.Kerberos.ClientPrincipal != "alice@CLIENT.EXAMPLE.COM" { + t.Fatalf("unexpected client principal: %q", config.Kerberos.ClientPrincipal) + } + if config.Kerberos.Realm != "CLIENT.EXAMPLE.COM" || config.Kerberos.CanonicalHostname { + t.Fatalf("unexpected Kerberos options: %#v", config.Kerberos) + } +} + +func TestKerberosReadsLegacyJAASKeytabOptions(t *testing.T) { + directory := t.TempDir() + configPath := filepath.Join(directory, "krb5.conf") + jaasPath := filepath.Join(directory, "hive jaas.conf") + keytabPath := filepath.Join(directory, "alice.keytab") + if err := os.WriteFile(configPath, []byte("[libdefaults]\n default_realm = EXAMPLE.COM\n"), 0o600); err != nil { + t.Fatal(err) + } + jaas := `HiveClient { + com.sun.security.auth.module.Krb5LoginModule required + useKeyTab=true + keyTab="` + keytabPath + `" + principal="alice@EXAMPLE.COM" + doNotPrompt=true; +};` + if err := os.WriteFile(jaasPath, []byte(jaas), 0o600); err != nil { + t.Fatal(err) + } + config, err := parseConnectionConfig(connectParams{ + ConnectionString: "jdbc:hive2://hs2.example.com/default;auth=kerberos;principal=hive/_HOST@EXAMPLE.COM", + AgentJavaOptions: []string{ + "-Djava.security.krb5.conf=" + configPath, + `-Djava.security.auth.login.config="` + jaasPath + `"`, + }, + }) + if err != nil { + t.Fatal(err) + } + if !config.Kerberos.UseKeytab || config.Kerberos.KeytabPath != keytabPath { + t.Fatalf("JAAS keytab was not preserved: %#v", config.Kerberos) + } + if config.Kerberos.ClientPrincipal != "alice@EXAMPLE.COM" { + t.Fatalf("JAAS principal was not preserved: %#v", config.Kerberos) + } +} + +func TestBuildHiveTLSConfigFromPKCS12Stores(t *testing.T) { + privateKey, certificate, _ := testZooKeeperCertificate(t) + password := "changeit" + keyStore, err := pkcs12.Modern.Encode(privateKey, certificate, nil, password) + if err != nil { + t.Fatal(err) + } + trustStore, err := pkcs12.Modern.EncodeTrustStore([]*x509.Certificate{certificate}, password) + if err != nil { + t.Fatal(err) + } + directory := t.TempDir() + keyPath := filepath.Join(directory, "client.p12") + trustPath := filepath.Join(directory, "trust.p12") + if err := os.WriteFile(keyPath, keyStore, 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(trustPath, trustStore, 0o600); err != nil { + t.Fatal(err) + } + config, err := buildTLSConfig(connectParams{}, map[string]string{ + "ssl": "true", + "twoway": "true", + "ssltruststore": trustPath, + "truststorepassword": password, + "truststoretype": "PKCS12", + "sslkeystore": keyPath, + "keystorepassword": password, + "keystoretype": "PKCS12", + }, "hs2.example.com") + if err != nil { + t.Fatal(err) + } + if config == nil || config.RootCAs == nil || len(config.Certificates) != 1 { + t.Fatalf("unexpected Hive TLS config: %#v", config) + } +} + +func TestBuildHiveTwoWayTLSRequiresTrustAndKeyMaterial(t *testing.T) { + if _, err := buildTLSConfig(connectParams{}, map[string]string{ + "ssl": "true", + "twoway": "true", + }, "hs2.example.com"); err == nil { + t.Fatal("expected missing two-way TLS material error") + } +} diff --git a/agents/drivers/argo-go/connector.go b/agents/drivers/argo-go/connector.go new file mode 100644 index 0000000000..ce745c1827 --- /dev/null +++ b/agents/drivers/argo-go/connector.go @@ -0,0 +1,199 @@ +package main + +import ( + "context" + "crypto/tls" + "database/sql" + "database/sql/driver" + "errors" + "fmt" + "strings" + "time" + + "github.com/beltran/gosasl" + gohive "github.com/t8y2/dbx/agents/go-common/gohive" +) + +type connectorFactory func(endpoint) driver.Connector + +type discoveryConnector struct { + discovery endpointDiscovery + factory connectorFactory + driver driver.Driver + retries int + retryInterval time.Duration +} + +func newDiscoveryConnector(config connectionConfig) *discoveryConnector { + return &discoveryConnector{ + discovery: newEndpointDiscovery(config), + factory: func(target endpoint) driver.Connector { + tlsConfig := config.TLSConfig + if target.SSL && tlsConfig == nil && !config.TLSExplicit { + tlsConfig = &tls.Config{MinVersion: tls.VersionTLS12, ServerName: target.Host} + } + if tlsConfig != nil { + tlsConfig = tlsConfig.Clone() + if tlsConfig.ServerName == "" || tlsConfig.ServerName == config.Endpoints[0].Host { + tlsConfig.ServerName = target.Host + } + } + hiveConfiguration := make(map[string]string, len(config.HiveConfiguration)+1) + for key, value := range config.HiveConfiguration { + hiveConfiguration[key] = value + } + if config.Kerberos.Enabled { + hiveConfiguration["hive.server2.thrift.sasl.qop"] = config.Kerberos.QOP + } + transportMode := config.TransportMode + if target.TransportMode != "" && !config.TransportModeExplicit { + transportMode = target.TransportMode + } + httpPath := config.HTTPPath + if target.HTTPPath != "" && !config.HTTPPathExplicit { + httpPath = target.HTTPPath + } + auth := config.Auth + if target.Auth != "" && !config.AuthExplicit { + auth = target.Auth + } + service := kerberosServiceForEndpoint(config, target) + return gohive.NewConnector(gohive.Config{ + Host: target.Host, + Port: target.Port, + Auth: normalizeHiveAuth(auth), + Username: config.Username, + Password: config.Password, + Database: config.Database, + TransportMode: transportMode, + HTTPPath: httpPath, + Service: service, + HTTPKerberosChannelBinding: config.Kerberos.ChannelBinding, + GSSAPIOptions: gssapiOptionsFromKerberos(config.Kerberos), + TLSConfig: tlsConfig, + HiveConfiguration: hiveConfiguration, + ConnectTimeout: config.ConnectTimeout, + SocketTimeout: config.SocketTimeout, + HTTPTimeout: config.SocketTimeout, + FetchSize: int64(config.FetchSize), + MaxMessageSize: config.MaxMessageSize, + HTTPHeaders: config.HTTPHeaders, + HTTPCookies: config.HTTPCookies, + RequestTracking: config.RequestTracking, + DisableCookieAuth: !config.CookieAuth, + CookieName: config.CookieName, + JWT: config.JWT, + DelegationToken: config.DelegationToken, + BrowserToken: config.BrowserToken, + BrowserClientID: config.BrowserClientID, + BrowserResponsePort: config.BrowserResponsePort, + BrowserResponseTimeout: config.BrowserResponseTimeout, + BrowserDisableSSLCheck: config.BrowserDisableSSLCheck, + WaitForNonQueryCompletion: waitsForNonQueryCompletion(config.DatabaseType), + }) + }, + driver: &gohive.Driver{}, + retries: max(config.Retries, 1), + retryInterval: config.RetryInterval, + } +} + +func waitsForNonQueryCompletion(databaseType string) bool { + return strings.EqualFold(databaseType, "impala") || strings.EqualFold(databaseType, "kyuubi") +} + +func gssapiOptionsFromKerberos(config kerberosConfig) gosasl.GSSAPIOptions { + return gosasl.GSSAPIOptions{ + ConfigPath: config.ConfigPath, + CCachePath: config.CCachePath, + KeytabPath: config.KeytabPath, + Principal: config.ClientPrincipal, + Password: config.Password, + QOP: config.QOP, + AuthorizationID: config.AuthorizationID, + ServiceHost: config.ServerName, + UseCCache: config.UseTicketCache, + UseKeytab: config.UseKeytab, + UseSSPI: config.UseSSPI, + CanonicalizeHost: config.CanonicalHostname, + DisablePAFXFAST: config.DisablePAFXFAST, + } +} + +func kerberosServiceForEndpoint(config connectionConfig, target endpoint) string { + service := firstNonEmpty(config.Kerberos.ServerPrincipal, config.Kerberos.Service) + if target.Principal != "" && !config.Kerberos.ServerPrincipalExplicit { + service = target.Principal + } + return service +} + +func (connector *discoveryConnector) Connect(ctx context.Context) (driver.Conn, error) { + var failures []string + for attempt := 0; attempt < max(connector.retries, 1); attempt++ { + rejected := map[string]bool{} + for { + endpoints, err := connector.discovery.Endpoints(ctx, rejected) + if err != nil { + failures = append(failures, fmt.Sprintf("discovery attempt %d: %v", attempt+1, err)) + break + } + if len(endpoints) == 0 { + break + } + for _, target := range endpoints { + connection, connectErr := connector.factory(target).Connect(ctx) + if connectErr == nil { + return connection, nil + } + rejected[target.address()] = true + failures = append(failures, fmt.Sprintf("attempt %d %s: %v", attempt+1, target.address(), connectErr)) + } + break + } + if attempt+1 < connector.retries && connector.retryInterval > 0 { + timer := time.NewTimer(connector.retryInterval) + select { + case <-ctx.Done(): + timer.Stop() + return nil, ctx.Err() + case <-timer.C: + } + } + } + if len(failures) == 0 { + return nil, errors.New("Hive discovery returned no endpoints") + } + return nil, fmt.Errorf("all HiveServer2 endpoints failed: %s", strings.Join(failures, "; ")) +} + +func (connector *discoveryConnector) Driver() driver.Driver { + return connector.driver +} + +func openHiveDatabase(config connectionConfig) *sql.DB { + database := sql.OpenDB(newDiscoveryConnector(config)) + database.SetMaxOpenConns(1) + database.SetMaxIdleConns(1) + return database +} + +func normalizeHiveAuth(value string) string { + normalized := strings.ToUpper(strings.TrimSpace(value)) + switch normalized { + case "": + return "NONE" + case "NOSASL", "NO_SASL": + return "NOSASL" + case "KERBEROS", "GSSAPI": + return "KERBEROS" + case "LDAP": + return "LDAP" + case "CUSTOM": + return "CUSTOM" + case "DIGEST-MD5", "DELEGATIONTOKEN", "DELEGATION_TOKEN": + return "DIGEST-MD5" + default: + return normalized + } +} diff --git a/agents/drivers/argo-go/connector_test.go b/agents/drivers/argo-go/connector_test.go new file mode 100644 index 0000000000..af9678152a --- /dev/null +++ b/agents/drivers/argo-go/connector_test.go @@ -0,0 +1,164 @@ +package main + +import ( + "context" + "database/sql/driver" + "errors" + "io" + "reflect" + "testing" +) + +type staticDiscovery struct { + values []endpoint +} + +func (discovery staticDiscovery) Endpoints(_ context.Context, rejected map[string]bool) ([]endpoint, error) { + result := make([]endpoint, 0, len(discovery.values)) + for _, value := range discovery.values { + if !rejected[value.address()] { + result = append(result, value) + } + } + return result, nil +} + +type fakeConnector struct { + connection driver.Conn + err error +} + +func (connector fakeConnector) Connect(context.Context) (driver.Conn, error) { + return connector.connection, connector.err +} + +func (fakeConnector) Driver() driver.Driver { return fakeDriver{} } + +type fakeDriver struct{} + +func (fakeDriver) Open(string) (driver.Conn, error) { return &fakeConnection{}, nil } + +type fakeConnection struct{} + +func (*fakeConnection) Prepare(string) (driver.Stmt, error) { return nil, errors.New("unsupported") } +func (*fakeConnection) Close() error { return nil } +func (*fakeConnection) Begin() (driver.Tx, error) { return nil, errors.New("unsupported") } + +type emptyRows struct{} + +func (emptyRows) Columns() []string { return nil } +func (emptyRows) Close() error { return nil } +func (emptyRows) Next([]driver.Value) error { return io.EOF } + +func TestDiscoveryConnectorFailsOver(t *testing.T) { + first := endpoint{Host: "first", Port: 10000} + second := endpoint{Host: "second", Port: 10000} + connected := &fakeConnection{} + var attempts []endpoint + connector := &discoveryConnector{ + discovery: staticDiscovery{values: []endpoint{first, second}}, + driver: fakeDriver{}, + factory: func(value endpoint) driver.Connector { + attempts = append(attempts, value) + if value == first { + return fakeConnector{err: errors.New("unavailable")} + } + return fakeConnector{connection: connected} + }, + } + value, err := connector.Connect(context.Background()) + if err != nil { + t.Fatal(err) + } + if value != connected || !reflect.DeepEqual(attempts, []endpoint{first, second}) { + t.Fatalf("unexpected failover result: value=%#v attempts=%#v", value, attempts) + } +} + +func TestDiscoveryConnectorRetriesAllEndpoints(t *testing.T) { + target := endpoint{Host: "hs2", Port: 10000} + connected := &fakeConnection{} + attempts := 0 + connector := &discoveryConnector{ + discovery: staticDiscovery{values: []endpoint{target}}, + driver: fakeDriver{}, + retries: 2, + factory: func(endpoint) driver.Connector { + attempts++ + if attempts == 1 { + return fakeConnector{err: errors.New("temporarily unavailable")} + } + return fakeConnector{connection: connected} + }, + } + value, err := connector.Connect(context.Background()) + if err != nil { + t.Fatal(err) + } + if value != connected || attempts != 2 { + t.Fatalf("unexpected retry result: value=%#v attempts=%d", value, attempts) + } +} + +func TestNormalizeHiveAuth(t *testing.T) { + for input, expected := range map[string]string{ + "": "NONE", + "noSasl": "NOSASL", + "kerberos": "KERBEROS", + "delegationToken": "DIGEST-MD5", + "vendor-auth-mode": "VENDOR-AUTH-MODE", + } { + if actual := normalizeHiveAuth(input); actual != expected { + t.Fatalf("normalizeHiveAuth(%q) = %q, expected %q", input, actual, expected) + } + } +} + +func TestWaitsForKyuubiAndImpalaNonQueryCompletion(t *testing.T) { + if !waitsForNonQueryCompletion("kyuubi") || !waitsForNonQueryCompletion("KYUUBI") || !waitsForNonQueryCompletion("impala") { + t.Fatal("Kyuubi and Impala must wait for asynchronous non-query operations") + } + if waitsForNonQueryCompletion("hive") { + t.Fatal("Hive must retain its existing non-query behavior") + } +} + +func TestKerberosServiceForEndpoint(t *testing.T) { + target := endpoint{Host: "hs2.example.com", Port: 10000, Principal: "hive/hs2.example.com@EXAMPLE.COM"} + discovered := connectionConfig{Kerberos: kerberosConfig{Service: "hive"}} + if value := kerberosServiceForEndpoint(discovered, target); value != target.Principal { + t.Fatalf("discovered principal was ignored: %s", value) + } + explicit := connectionConfig{Kerberos: kerberosConfig{ + Service: "hive", + ServerPrincipal: "hive/_HOST@USER.EXAMPLE.COM", + ServerPrincipalExplicit: true, + }} + if value := kerberosServiceForEndpoint(explicit, target); value != explicit.Kerberos.ServerPrincipal { + t.Fatalf("explicit principal was overwritten: %s", value) + } +} + +func TestKerberosUsesConnectionScopedGSSAPIOptions(t *testing.T) { + config := kerberosConfig{ + Enabled: true, + ServerPrincipal: "hive/_HOST@EXAMPLE.COM", + ClientPrincipal: "alice@EXAMPLE.COM", + ServerName: "canonical.example.com", + CanonicalHostname: true, + ConfigPath: "/etc/krb5.conf", + CCachePath: "/tmp/alice.ccache", + KeytabPath: "/tmp/alice.keytab", + Password: "secret", + AuthorizationID: "proxy-user", + QOP: "auth-conf", + UseTicketCache: true, + UseKeytab: true, + UseSSPI: true, + DisablePAFXFAST: true, + } + options := gssapiOptionsFromKerberos(config) + if options.ConfigPath != "/etc/krb5.conf" || options.Principal != "alice@EXAMPLE.COM" || options.Password != "secret" || options.AuthorizationID != "proxy-user" || options.ServiceHost != "canonical.example.com" || !options.CanonicalizeHost || options.CCachePath != "/tmp/alice.ccache" || options.KeytabPath != "/tmp/alice.keytab" || options.QOP != "auth-conf" || !options.UseCCache || !options.UseKeytab || !options.UseSSPI || !options.DisablePAFXFAST { + t.Fatalf("unexpected GSSAPI options: %#v", options) + } +} diff --git a/agents/drivers/argo-go/discovery.go b/agents/drivers/argo-go/discovery.go new file mode 100644 index 0000000000..94e01f156c --- /dev/null +++ b/agents/drivers/argo-go/discovery.go @@ -0,0 +1,392 @@ +package main + +import ( + "context" + "crypto/tls" + "encoding/json" + "errors" + "fmt" + "math/rand/v2" + "net" + "net/url" + "strconv" + "strings" + "time" + + "github.com/go-zookeeper/zk" +) + +type endpointDiscovery interface { + Endpoints(context.Context, map[string]bool) ([]endpoint, error) +} + +type directDiscovery struct { + endpoints []endpoint +} + +func (discovery directDiscovery) Endpoints(_ context.Context, rejected map[string]bool) ([]endpoint, error) { + return shuffledEndpoints(discovery.endpoints, rejected), nil +} + +type zooKeeperDiscovery struct { + servers []endpoint + namespace string + discoveryMode string + authScheme string + auth string + timeout time.Duration + dialer func([]string, time.Duration) (zooKeeperClient, <-chan zk.Event, error) +} + +type zooKeeperClient interface { + AddAuth(string, []byte) error + Children(string) ([]string, *zk.Stat, error) + Get(string) ([]byte, *zk.Stat, error) + Close() +} + +func newEndpointDiscovery(config connectionConfig) endpointDiscovery { + if strings.EqualFold(config.ServiceDiscoveryMode, "zookeeper") || strings.EqualFold(config.ServiceDiscoveryMode, "zookeeperha") { + return &zooKeeperDiscovery{ + servers: append([]endpoint(nil), config.Endpoints...), + namespace: config.ZooKeeperNamespace, + discoveryMode: config.ServiceDiscoveryMode, + authScheme: config.ZooKeeperAuthScheme, + auth: config.ZooKeeperAuth, + timeout: config.ConnectTimeout, + dialer: newZooKeeperDialer(config), + } + } + return directDiscovery{endpoints: append([]endpoint(nil), config.Endpoints...)} +} + +func newZooKeeperDialer(config connectionConfig) func([]string, time.Duration) (zooKeeperClient, <-chan zk.Event, error) { + if config.ZooKeeperKerberos.Enabled { + return func(servers []string, timeout time.Duration) (zooKeeperClient, <-chan zk.Event, error) { + return connectKerberosZooKeeper(servers, timeout, config.ZooKeeperTLSConfig, config) + } + } + tlsConfig := config.ZooKeeperTLSConfig + if tlsConfig == nil { + return func(servers []string, timeout time.Duration) (zooKeeperClient, <-chan zk.Event, error) { + return zk.Connect(servers, timeout, zk.WithLogInfo(false)) + } + } + return func(servers []string, timeout time.Duration) (zooKeeperClient, <-chan zk.Event, error) { + return zk.Connect( + servers, + timeout, + zk.WithLogInfo(false), + zk.WithDialer(func(network, address string, dialTimeout time.Duration) (net.Conn, error) { + config := tlsConfig.Clone() + if config.ServerName == "" { + host, _, err := net.SplitHostPort(address) + if err == nil { + config.ServerName = host + } + } + dialer := &net.Dialer{Timeout: dialTimeout} + return tls.DialWithDialer(dialer, network, address, config) + }), + ) + } +} + +func (discovery *zooKeeperDiscovery) Endpoints(ctx context.Context, rejected map[string]bool) ([]endpoint, error) { + addresses := make([]string, 0, len(discovery.servers)) + for _, server := range discovery.servers { + addresses = append(addresses, server.address()) + } + timeout := discovery.timeout + if timeout <= 0 { + timeout = defaultConnectTimeout + } + connection, events, err := discovery.dialer(addresses, timeout) + if err != nil { + return nil, fmt.Errorf("connect to ZooKeeper: %w", err) + } + defer connection.Close() + if err := waitForZooKeeperSession(ctx, events, timeout); err != nil { + return nil, err + } + if discovery.authScheme != "" || discovery.auth != "" { + if discovery.authScheme == "" || discovery.auth == "" { + return nil, errors.New("ZooKeeper auth scheme and credentials must be configured together") + } + if err := connection.AddAuth(discovery.authScheme, []byte(discovery.auth)); err != nil { + return nil, fmt.Errorf("authenticate to ZooKeeper: %w", err) + } + } + resolved := make([]endpoint, 0) + var listedPath string + var nodeFailures []string + for _, path := range discovery.paths() { + children, _, childrenErr := connection.Children(path) + if errors.Is(childrenErr, zk.ErrNoNode) { + continue + } + if childrenErr != nil { + return nil, fmt.Errorf("list ZooKeeper namespace %s: %w", path, childrenErr) + } + listedPath = path + for _, child := range children { + data, _, dataErr := connection.Get(path + "/" + child) + if dataErr != nil { + if errors.Is(dataErr, zk.ErrNoNode) { + continue + } + nodeFailures = append(nodeFailures, fmt.Sprintf("%s/%s: %v", path, child, dataErr)) + continue + } + value, parseErr := parseHiveServerRegistration(child, data) + if parseErr == nil { + resolved = append(resolved, value) + } else { + nodeFailures = append(nodeFailures, fmt.Sprintf("%s/%s: %v", path, child, parseErr)) + } + } + if len(resolved) > 0 { + break + } + } + resolved = shuffledEndpoints(uniqueEndpoints(resolved), rejected) + if len(resolved) == 0 { + if listedPath == "" { + return nil, fmt.Errorf("HiveServer2 ZooKeeper namespace not found; tried %s", strings.Join(discovery.paths(), ", ")) + } + if len(nodeFailures) > 0 { + return nil, fmt.Errorf("no usable HiveServer2 nodes in ZooKeeper namespace %s: %s", listedPath, strings.Join(nodeFailures, "; ")) + } + return nil, fmt.Errorf("no available HiveServer2 nodes in ZooKeeper namespace %s", listedPath) + } + return resolved, nil +} + +func (discovery *zooKeeperDiscovery) paths() []string { + namespace := strings.Trim(discovery.namespace, "/") + if strings.EqualFold(discovery.discoveryMode, "zookeeperha") { + return []string{ + zooKeeperPath(namespace, "instances"), + zooKeeperPath(namespace+"-unsecure", "instances"), + zooKeeperPath(namespace+"-sasl", "instances"), + } + } + return []string{zooKeeperPath(namespace)} +} + +func zooKeeperPath(parts ...string) string { + cleaned := make([]string, 0, len(parts)) + for _, part := range parts { + if value := strings.Trim(part, "/"); value != "" { + cleaned = append(cleaned, value) + } + } + if len(cleaned) == 0 { + return "/" + } + return "/" + strings.Join(cleaned, "/") +} + +func waitForZooKeeperSession(ctx context.Context, events <-chan zk.Event, timeout time.Duration) error { + timer := time.NewTimer(timeout) + defer timer.Stop() + for { + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return errors.New("ZooKeeper connection timed out before a session was established") + case event, ok := <-events: + if !ok { + return errors.New("ZooKeeper event stream closed before a session was established") + } + if event.Err != nil { + return fmt.Errorf("ZooKeeper connection event: %w", event.Err) + } + switch event.State { + case zk.StateHasSession: + return nil + case zk.StateAuthFailed: + return errors.New("ZooKeeper authentication failed") + case zk.StateExpired: + return errors.New("ZooKeeper session expired during connection") + } + } + } +} + +func parseHiveServerRegistration(child string, data []byte) (endpoint, error) { + candidates := []string{strings.TrimSpace(string(data)), strings.TrimSpace(child)} + for _, candidate := range candidates { + if candidate == "" { + continue + } + if value, err := endpointFromRegistrationJSON(candidate); err == nil { + return value, nil + } + parameters := parseHiveParameters(candidate) + for _, key := range []string{"serveruri", "hiveserver2uri", "server_uri"} { + if raw := parameter(parameters, key); raw != "" { + return parseRegisteredEndpoint(raw) + } + } + if value, err := endpointFromPublishedHiveConfig(parameters); err == nil { + return value, nil + } + if strings.Contains(candidate, "=") { + continue + } + if value, err := parseRegisteredEndpoint(candidate); err == nil { + return value, nil + } + } + return endpoint{}, fmt.Errorf("unsupported HiveServer2 ZooKeeper registration %q", child) +} + +func endpointFromRegistrationJSON(value string) (endpoint, error) { + var object map[string]any + if json.Unmarshal([]byte(value), &object) != nil { + return endpoint{}, errors.New("not JSON") + } + for _, key := range []string{"serverUri", "server_uri", "hiveServer2Uri", "uri"} { + if raw, ok := object[key].(string); ok && strings.TrimSpace(raw) != "" { + return parseRegisteredEndpoint(raw) + } + } + if serviceRecordEndpoint, ok := endpointFromServiceRecord(object); ok { + return serviceRecordEndpoint, nil + } + host, _ := object["host"].(string) + if host == "" { + host, _ = object["hostname"].(string) + } + port := 0 + switch value := object["port"].(type) { + case float64: + port = int(value) + case string: + port, _ = strconv.Atoi(value) + } + if host != "" && port > 0 { + return endpoint{Host: host, Port: port}, nil + } + return endpoint{}, errors.New("JSON registration has no endpoint") +} + +func endpointFromServiceRecord(object map[string]any) (endpoint, bool) { + internal, _ := object["internal"].([]any) + for _, rawEndpoint := range internal { + published, _ := rawEndpoint.(map[string]any) + if !strings.EqualFold(registrationStringValue(published["api"]), "activeEndpoint") { + continue + } + addresses, _ := published["addresses"].([]any) + for _, rawAddress := range addresses { + address, _ := rawAddress.(map[string]any) + host := registrationStringValue(address["host"]) + port, _ := strconv.Atoi(registrationStringValue(address["port"])) + if host != "" && port > 0 { + result := endpoint{Host: host, Port: port} + applyPublishedHiveConfig(&result, object) + return result, true + } + } + } + return endpoint{}, false +} + +func registrationStringValue(value any) string { + switch typed := value.(type) { + case string: + return strings.TrimSpace(typed) + case float64: + return strconv.FormatInt(int64(typed), 10) + case json.Number: + return typed.String() + default: + return "" + } +} + +func applyPublishedHiveConfig(target *endpoint, parameters map[string]any) { + value := func(key string) string { + if nested, ok := parameters["attributes"].(map[string]any); ok { + if result := registrationStringValue(nested[key]); result != "" { + return result + } + } + return registrationStringValue(parameters[key]) + } + target.TransportMode = strings.ToLower(value("hive.server2.transport.mode")) + target.HTTPPath = strings.TrimPrefix(value("hive.server2.thrift.http.path"), "/") + target.Auth = strings.ToUpper(value("hive.server2.authentication")) + target.Principal = value("hive.server2.authentication.kerberos.principal") + target.SSL = strings.EqualFold(value("hive.server2.use.ssl"), "true") +} + +func endpointFromPublishedHiveConfig(parameters map[string]string) (endpoint, error) { + host := firstNonEmpty( + parameter(parameters, "hive.server2.thrift.bind.host"), + parameter(parameters, "host"), + ) + transportMode := strings.ToLower(parameter(parameters, "hive.server2.transport.mode")) + portValue := parameter(parameters, "hive.server2.thrift.port") + if transportMode == "http" { + portValue = firstNonEmpty(parameter(parameters, "hive.server2.thrift.http.port"), portValue) + } + port, err := strconv.Atoi(portValue) + if host == "" || err != nil || port <= 0 { + return endpoint{}, errors.New("published HiveServer2 configuration has no valid host and port") + } + return endpoint{ + Host: host, + Port: port, + TransportMode: transportMode, + HTTPPath: strings.TrimPrefix(parameter(parameters, "hive.server2.thrift.http.path"), "/"), + Auth: strings.ToUpper(parameter(parameters, "hive.server2.authentication")), + Principal: parameter(parameters, "hive.server2.authentication.kerberos.principal"), + SSL: parameterBool(parameters, "hive.server2.use.ssl"), + }, nil +} + +func parseRegisteredEndpoint(value string) (endpoint, error) { + value = strings.TrimSpace(value) + if parsed, err := url.Parse(value); err == nil && parsed.Hostname() != "" { + port := defaultHivePort + if parsed.Port() != "" { + parsedPort, parseErr := strconv.Atoi(parsed.Port()) + if parseErr != nil { + return endpoint{}, parseErr + } + port = parsedPort + } + return endpoint{Host: parsed.Hostname(), Port: port}, nil + } + return parseEndpoint(value, defaultHivePort) +} + +func shuffledEndpoints(values []endpoint, rejected map[string]bool) []endpoint { + result := make([]endpoint, 0, len(values)) + for _, value := range values { + if !rejected[value.address()] { + result = append(result, value) + } + } + rand.Shuffle(len(result), func(first, second int) { + result[first], result[second] = result[second], result[first] + }) + return result +} + +func uniqueEndpoints(values []endpoint) []endpoint { + seen := map[string]bool{} + result := make([]endpoint, 0, len(values)) + for _, value := range values { + key := value.address() + if !seen[key] { + seen[key] = true + result = append(result, value) + } + } + return result +} diff --git a/agents/drivers/argo-go/discovery_test.go b/agents/drivers/argo-go/discovery_test.go new file mode 100644 index 0000000000..dfa14f9e43 --- /dev/null +++ b/agents/drivers/argo-go/discovery_test.go @@ -0,0 +1,221 @@ +package main + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/go-zookeeper/zk" +) + +func TestParseHiveServerRegistrationFromChildName(t *testing.T) { + value, err := parseHiveServerRegistration("serverUri=hs2.example.com:10000;version=4.0.1", nil) + if err != nil { + t.Fatal(err) + } + if value != (endpoint{Host: "hs2.example.com", Port: 10000}) { + t.Fatalf("unexpected endpoint: %#v", value) + } +} + +func TestParseHiveServerRegistrationFromData(t *testing.T) { + value, err := parseHiveServerRegistration("instance-0001", []byte(`{"serverUri":"thrift://kyuubi.example.com:10009"}`)) + if err != nil { + t.Fatal(err) + } + if value != (endpoint{Host: "kyuubi.example.com", Port: 10009}) { + t.Fatalf("unexpected endpoint: %#v", value) + } +} + +func TestRejectedEndpointsAreExcluded(t *testing.T) { + values := shuffledEndpoints([]endpoint{{Host: "one", Port: 1}, {Host: "two", Port: 2}}, map[string]bool{"one:1": true}) + if len(values) != 1 || values[0].Host != "two" { + t.Fatalf("unexpected endpoints: %#v", values) + } +} + +func TestParsePublishedHiveServerConfiguration(t *testing.T) { + value, err := parseHiveServerRegistration("instance-1", []byte( + "hive.server2.thrift.bind.host=hs2.example.com;"+ + "hive.server2.transport.mode=http;"+ + "hive.server2.thrift.http.port=10001;"+ + "hive.server2.thrift.http.path=cliservice;"+ + "hive.server2.authentication=KERBEROS;"+ + "hive.server2.authentication.kerberos.principal=hive/_HOST@EXAMPLE.COM;"+ + "hive.server2.use.ssl=true", + )) + if err != nil { + t.Fatal(err) + } + if value.Host != "hs2.example.com" || value.Port != 10001 || value.TransportMode != "http" || value.HTTPPath != "cliservice" || value.Auth != "KERBEROS" || !value.SSL { + t.Fatalf("unexpected endpoint: %#v", value) + } +} + +func TestParseActivePassiveServiceRecord(t *testing.T) { + value, err := parseHiveServerRegistration("instance-0001", []byte(`{ + "hive.server2.transport.mode":"binary", + "hive.server2.authentication":"KERBEROS", + "hive.server2.authentication.kerberos.principal":"hive/_HOST@EXAMPLE.COM", + "internal":[{ + "api":"activeEndpoint", + "addresses":[{"host":"active.example.com","port":"10000"}] + }] + }`)) + if err != nil { + t.Fatal(err) + } + if value.Host != "active.example.com" || value.Port != 10000 || value.Auth != "KERBEROS" || value.TransportMode != "binary" { + t.Fatalf("unexpected active endpoint: %#v", value) + } +} + +func TestZooKeeperHAPaths(t *testing.T) { + discovery := &zooKeeperDiscovery{namespace: "hs2ActivePassiveHA", discoveryMode: "zooKeeperHA"} + want := []string{ + "/hs2ActivePassiveHA/instances", + "/hs2ActivePassiveHA-unsecure/instances", + "/hs2ActivePassiveHA-sasl/instances", + } + got := discovery.paths() + if len(got) != len(want) { + t.Fatalf("unexpected paths: %#v", got) + } + for index := range want { + if got[index] != want[index] { + t.Fatalf("path %d = %q, want %q", index, got[index], want[index]) + } + } +} + +func TestZooKeeperPathNormalizesEmptyNamespace(t *testing.T) { + if value := zooKeeperPath(""); value != "/" { + t.Fatalf("unexpected root path: %q", value) + } + if value := zooKeeperPath("", "instances"); value != "/instances" { + t.Fatalf("unexpected instances path: %q", value) + } +} + +func TestWaitForZooKeeperSession(t *testing.T) { + events := make(chan zk.Event, 2) + events <- zk.Event{State: zk.StateConnected} + events <- zk.Event{State: zk.StateHasSession} + if err := waitForZooKeeperSession(context.Background(), events, time.Second); err != nil { + t.Fatal(err) + } +} + +func TestWaitForZooKeeperSessionRejectsAuthFailure(t *testing.T) { + events := make(chan zk.Event, 1) + events <- zk.Event{State: zk.StateAuthFailed} + if err := waitForZooKeeperSession(context.Background(), events, time.Second); err == nil { + t.Fatal("expected authentication failure") + } +} + +type fakeZooKeeperClient struct { + children map[string][]string + data map[string][]byte + auth []byte + closed bool +} + +func (client *fakeZooKeeperClient) AddAuth(_ string, auth []byte) error { + client.auth = append([]byte(nil), auth...) + return nil +} + +func (client *fakeZooKeeperClient) Children(path string) ([]string, *zk.Stat, error) { + children, ok := client.children[path] + if !ok { + return nil, nil, zk.ErrNoNode + } + return append([]string(nil), children...), nil, nil +} + +func (client *fakeZooKeeperClient) Get(path string) ([]byte, *zk.Stat, error) { + data, ok := client.data[path] + if !ok { + return nil, nil, zk.ErrNoNode + } + return append([]byte(nil), data...), nil, nil +} + +func (client *fakeZooKeeperClient) Close() { client.closed = true } + +func TestZooKeeperDiscoveryUsesSessionAuthAndSkipsStaleNodes(t *testing.T) { + client := &fakeZooKeeperClient{ + children: map[string][]string{"/hiveserver2": {"stale", "live"}}, + data: map[string][]byte{ + "/hiveserver2/live": []byte("serverUri=live.example.com:10000"), + }, + } + discovery := &zooKeeperDiscovery{ + servers: []endpoint{{Host: "zk.example.com", Port: 2181}}, + namespace: "hiveserver2", + authScheme: "digest", + auth: "user:password", + timeout: time.Second, + discoveryMode: "zookeeper", + dialer: func([]string, time.Duration) (zooKeeperClient, <-chan zk.Event, error) { + events := make(chan zk.Event, 1) + events <- zk.Event{State: zk.StateHasSession} + return client, events, nil + }, + } + values, err := discovery.Endpoints(context.Background(), map[string]bool{}) + if err != nil { + t.Fatal(err) + } + if len(values) != 1 || values[0].Host != "live.example.com" { + t.Fatalf("unexpected endpoints: %#v", values) + } + if string(client.auth) != "user:password" || !client.closed { + t.Fatalf("auth=%q closed=%v", client.auth, client.closed) + } +} + +func TestZooKeeperHAFallsThroughStaleCandidatePath(t *testing.T) { + client := &fakeZooKeeperClient{ + children: map[string][]string{ + "/hs2ActivePassiveHA/instances": {"stale"}, + "/hs2ActivePassiveHA-unsecure/instances": {"live"}, + }, + data: map[string][]byte{ + "/hs2ActivePassiveHA-unsecure/instances/live": []byte("serverUri=active.example.com:10000"), + }, + } + discovery := &zooKeeperDiscovery{ + servers: []endpoint{{Host: "zk.example.com", Port: 2181}}, + namespace: "hs2ActivePassiveHA", + timeout: time.Second, + discoveryMode: "zookeeperHA", + dialer: func([]string, time.Duration) (zooKeeperClient, <-chan zk.Event, error) { + events := make(chan zk.Event, 1) + events <- zk.Event{State: zk.StateHasSession} + return client, events, nil + }, + } + values, err := discovery.Endpoints(context.Background(), map[string]bool{}) + if err != nil { + t.Fatal(err) + } + if len(values) != 1 || values[0].Host != "active.example.com" { + t.Fatalf("unexpected HA endpoints: %#v", values) + } +} + +func TestZooKeeperDiscoveryPropagatesDialError(t *testing.T) { + discovery := &zooKeeperDiscovery{ + servers: []endpoint{{Host: "zk.example.com", Port: 2181}}, + dialer: func([]string, time.Duration) (zooKeeperClient, <-chan zk.Event, error) { + return nil, nil, errors.New("dial failed") + }, + } + if _, err := discovery.Endpoints(context.Background(), nil); err == nil { + t.Fatal("expected dial error") + } +} diff --git a/agents/drivers/argo-go/go.mod b/agents/drivers/argo-go/go.mod new file mode 100644 index 0000000000..54d85fc8d6 --- /dev/null +++ b/agents/drivers/argo-go/go.mod @@ -0,0 +1,35 @@ +module github.com/t8y2/dbx/agents/drivers/argo-go + +go 1.23.0 + +require ( + github.com/beltran/gosasl v1.0.0 + github.com/go-zookeeper/zk v1.0.4 + github.com/golang-auth/go-gssapi/v2 v2.0.0 + github.com/jcmturner/krb5test v0.0.0-20201230140143-102e4b78cdb8 + github.com/pavlo-v-chernykh/keystore-go/v4 v4.5.0 + github.com/t8y2/dbx/agents/go-common/gohive v0.0.0 + software.sslmate.com/src/go-pkcs12 v0.7.3 +) + +require ( + github.com/alexbrainman/sspi v0.0.0-20250919150558-7d374ff0d59e // indirect + github.com/apache/thrift v0.22.0 // indirect + github.com/beltran/gohive/v2 v2.1.0 // indirect + github.com/hashicorp/go-uuid v1.0.2 // indirect + github.com/jcmturner/aescts/v2 v2.0.0 // indirect + github.com/jcmturner/dnsutils/v2 v2.0.0 // indirect + github.com/jcmturner/gofork v1.0.0 // indirect + github.com/jcmturner/gokrb5 v8.4.2+incompatible // indirect + github.com/jcmturner/gokrb5/v8 v8.4.2 // indirect + github.com/jcmturner/rpc/v2 v2.0.3 // indirect + github.com/pkg/errors v0.9.1 // indirect + golang.org/x/crypto v0.39.0 // indirect + golang.org/x/net v0.41.0 // indirect +) + +replace github.com/beltran/gosasl => ../../go-common/gosasl + +replace github.com/golang-auth/go-gssapi/v2 => ../../go-common/go-gssapi + +replace github.com/t8y2/dbx/agents/go-common/gohive => ../../go-common/gohive diff --git a/agents/drivers/argo-go/go.sum b/agents/drivers/argo-go/go.sum new file mode 100644 index 0000000000..04407f477b --- /dev/null +++ b/agents/drivers/argo-go/go.sum @@ -0,0 +1,68 @@ +github.com/alexbrainman/sspi v0.0.0-20250919150558-7d374ff0d59e h1:4dAU9FXIyQktpoUAgOJK3OTFc/xug0PCXYCqU0FgDKI= +github.com/alexbrainman/sspi v0.0.0-20250919150558-7d374ff0d59e/go.mod h1:cEWa1LVoE5KvSD9ONXsZrj0z6KqySlCCNKHlLzbqAt4= +github.com/apache/thrift v0.22.0 h1:r7mTJdj51TMDe6RtcmNdQxgn9XcyfGDOzegMDRg47uc= +github.com/apache/thrift v0.22.0/go.mod h1:1e7J/O1Ae6ZQMTYdy9xa3w9k+XHWPfRvdPyJeynQ+/g= +github.com/beltran/gohive/v2 v2.1.0 h1:+XPfODaLXsZkD8mEnKCYqinC4OmVMx4/oR4ys8IGKeI= +github.com/beltran/gohive/v2 v2.1.0/go.mod h1:MrxyI0sdriG52I+Bo2lJpGR75bxjcI/Hx3+MBnYqR10= +github.com/davecgh/go-spew v1.1.0 h1:ZDRjVQ15GmhC3fiQ8ni8+OwkZQO4DARzQgrnXU1Liz8= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/go-zookeeper/zk v1.0.4 h1:DPzxraQx7OrPyXq2phlGlNSIyWEsAox0RJmjTseMV6I= +github.com/go-zookeeper/zk v1.0.4/go.mod h1:nOB03cncLtlp4t+UAkGSV+9beXP/akpekBwL+UX1Qcw= +github.com/gorilla/securecookie v1.1.1/go.mod h1:ra0sb63/xPlUeL+yeDciTfxMRAA+MP+HVt/4epWDjd4= +github.com/gorilla/sessions v1.2.0/go.mod h1:dk2InVEVJ0sfLlnXv9EAgkf6ecYs/i80K/zI+bUmuGM= +github.com/gorilla/sessions v1.2.1/go.mod h1:dk2InVEVJ0sfLlnXv9EAgkf6ecYs/i80K/zI+bUmuGM= +github.com/hashicorp/go-uuid v1.0.2 h1:cfejS+Tpcp13yd5nYHWDI6qVCny6wyX2Mt5SGur2IGE= +github.com/hashicorp/go-uuid v1.0.2/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro= +github.com/jcmturner/aescts/v2 v2.0.0 h1:9YKLH6ey7H4eDBXW8khjYslgyqG2xZikXP0EQFKrle8= +github.com/jcmturner/aescts/v2 v2.0.0/go.mod h1:AiaICIRyfYg35RUkr8yESTqvSy7csK90qZ5xfvvsoNs= +github.com/jcmturner/dnsutils/v2 v2.0.0 h1:lltnkeZGL0wILNvrNiVCR6Ro5PGU/SeBvVO/8c/iPbo= +github.com/jcmturner/dnsutils/v2 v2.0.0/go.mod h1:b0TnjGOvI/n42bZa+hmXL+kFJZsFT7G4t3HTlQ184QM= +github.com/jcmturner/gofork v1.0.0 h1:J7uCkflzTEhUZ64xqKnkDxq3kzc96ajM1Gli5ktUem8= +github.com/jcmturner/gofork v1.0.0/go.mod h1:MK8+TM0La+2rjBD4jE12Kj1pCCxK7d2LK/UM3ncEo0o= +github.com/jcmturner/goidentity/v6 v6.0.1 h1:VKnZd2oEIMorCTsFBnJWbExfNN7yZr3EhJAxwOkZg6o= +github.com/jcmturner/goidentity/v6 v6.0.1/go.mod h1:X1YW3bgtvwAXju7V3LCIMpY0Gbxyjn/mY9zx4tFonSg= +github.com/jcmturner/gokrb5 v8.4.2+incompatible h1:MQW70Fbazv31g6URAXCjO2bGenIL0wVt3wqcpc0EjHI= +github.com/jcmturner/gokrb5 v8.4.2+incompatible/go.mod h1:0Q5eFyVvYsEsZ8xl1A/jUqhXvxUp/X9ELrJm+zieq5E= +github.com/jcmturner/gokrb5/v8 v8.4.0/go.mod h1:T1hnNppQsBtxW0tCHMHTkAt8n/sABdzZgZdoFrZaZNM= +github.com/jcmturner/gokrb5/v8 v8.4.2 h1:6ZIM6b/JJN0X8UM43ZOM6Z4SJzla+a/u7scXFJzodkA= +github.com/jcmturner/gokrb5/v8 v8.4.2/go.mod h1:sb+Xq/fTY5yktf/VxLsE3wlfPqQjp0aWNYyvBVK62bc= +github.com/jcmturner/krb5test v0.0.0-20201230140143-102e4b78cdb8 h1:bWgHpg2hLK3m8zHGSLqa+UV8XkivFFjG5IniW5lWzIA= +github.com/jcmturner/krb5test v0.0.0-20201230140143-102e4b78cdb8/go.mod h1:eUQyrWMo2eMPWPezAYjLLwaiQHzJUD3su+1vAL10dOk= +github.com/jcmturner/rpc/v2 v2.0.2/go.mod h1:VUJYCIDm3PVOEHw8sgt091/20OJjskO/YJki3ELg/Hc= +github.com/jcmturner/rpc/v2 v2.0.3 h1:7FXXj8Ti1IaVFpSAziCZWNzbNuZmnvw/i6CqLNdWfZY= +github.com/jcmturner/rpc/v2 v2.0.3/go.mod h1:VUJYCIDm3PVOEHw8sgt091/20OJjskO/YJki3ELg/Hc= +github.com/pavlo-v-chernykh/keystore-go/v4 v4.5.0 h1:2nosf3P75OZv2/ZO/9Px5ZgZ5gbKrzA3joN1QMfOGMQ= +github.com/pavlo-v-chernykh/keystore-go/v4 v4.5.0/go.mod h1:lAVhWwbNaveeJmxrxuSTxMgKpF6DjnuVpn6T8WiBwYQ= +github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= +github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4= +github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.0 h1:nwc3DEeHmmLAfoZucVR881uASk0Mfjw8xYJ99tb5CcY= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= +golang.org/x/crypto v0.0.0-20200117160349-530e935923ad/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= +golang.org/x/crypto v0.0.0-20201112155050-0c6587e931a9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= +golang.org/x/crypto v0.39.0 h1:SHs+kF4LP+f+p14esP5jAoDpHU8Gu/v9lFRK6IT5imM= +golang.org/x/crypto v0.39.0/go.mod h1:L+Xg3Wf6HoL4Bn4238Z6ft6KfEpN0tJGo53AAPC632U= +golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/net v0.0.0-20200114155413-6afb5195e5aa/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= +golang.org/x/net v0.0.0-20210326220855-61e056675ecf/go.mod h1:uSPa2vr4CLtc/ILN5odXGNXS6mhrKVzTaCXzk9m6W3k= +golang.org/x/net v0.41.0 h1:vBTly1HeNPEn3wtREYfy4GZ/NECgw2Cnl+nK6Nz3uvw= +golang.org/x/net v0.41.0/go.mod h1:B/K4NNqkfmg07DQYrbwvSluqCJOOXwUjeb/5lOisjbA= +golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= +golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210324051608-47abb6519492/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= +golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= +golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +software.sslmate.com/src/go-pkcs12 v0.7.3 h1:JBQD3FDqYjTeyDAeZQklj2ar88ykBLtALloPJHyAauU= +software.sslmate.com/src/go-pkcs12 v0.7.3/go.mod h1:Qiz0EyvDRJjjxGyUQa2cCNZn/wMyzrRJ/qcDXOQazLI= diff --git a/agents/drivers/argo-go/init_test.go b/agents/drivers/argo-go/init_test.go new file mode 100644 index 0000000000..9ad3edd58d --- /dev/null +++ b/agents/drivers/argo-go/init_test.go @@ -0,0 +1,65 @@ +package main + +import ( + "context" + "database/sql/driver" + "errors" + "reflect" + "strings" + "testing" + + "github.com/t8y2/dbx/agents/go-common/gohive" +) + +func TestRunHiveInitStatementsExecutesAndDrainsResults(t *testing.T) { + var resultRows *scriptedRows + behavior := &scriptedBehavior{} + behavior.query = func(ctx context.Context, query string) (driver.Rows, error) { + switch query { + case "SET hive.exec.dynamic.partition=true": + resultRows = newScriptedRows(ctx, []string{"set"}, []string{"STRING"}, [][]driver.Value{ + {"hive.exec.dynamic.partition=true"}, + }) + return resultRows, nil + case "USE analytics": + return nil, &gohive.NonQueryResult{} + default: + return nil, errors.New("unexpected init statement") + } + } + server := newScriptedServer(t, behavior) + + err := runHiveInitStatements( + context.Background(), + server.connection, + []string{"SET hive.exec.dynamic.partition=true", "USE analytics"}, + 77, + ) + if err != nil { + t.Fatal(err) + } + queries, executions, _, _ := behavior.snapshot() + if !reflect.DeepEqual(queries, []string{"SET hive.exec.dynamic.partition=true", "USE analytics"}) || len(executions) != 0 { + t.Fatalf("unexpected init statements: queries=%v executions=%v", queries, executions) + } + if resultRows == nil || !resultRows.isClosed() { + t.Fatal("initFile result set was not drained and closed") + } +} + +func TestRunHiveInitStatementsStopsOnFailure(t *testing.T) { + behavior := &scriptedBehavior{} + behavior.query = func(context.Context, string) (driver.Rows, error) { + return nil, errors.New("permission denied") + } + server := newScriptedServer(t, behavior) + + err := runHiveInitStatements(context.Background(), server.connection, []string{"USE restricted", "USE skipped"}, 100) + if err == nil || !strings.Contains(err.Error(), "execute Hive initFile statement: permission denied") { + t.Fatalf("unexpected initFile error: %v", err) + } + queries, _, _, _ := behavior.snapshot() + if !reflect.DeepEqual(queries, []string{"USE restricted"}) { + t.Fatalf("initFile continued after failure: %v", queries) + } +} diff --git a/agents/drivers/argo-go/kerberos_defaults_unix.go b/agents/drivers/argo-go/kerberos_defaults_unix.go new file mode 100644 index 0000000000..6a6e890932 --- /dev/null +++ b/agents/drivers/argo-go/kerberos_defaults_unix.go @@ -0,0 +1,16 @@ +//go:build !windows + +package main + +import ( + "fmt" + "os" +) + +func defaultKerberosConfigPath() string { + return "/etc/krb5.conf" +} + +func defaultKerberosCCachePath() string { + return fmt.Sprintf("/tmp/krb5cc_%d", os.Getuid()) +} diff --git a/agents/drivers/argo-go/kerberos_defaults_windows.go b/agents/drivers/argo-go/kerberos_defaults_windows.go new file mode 100644 index 0000000000..1ed65fac33 --- /dev/null +++ b/agents/drivers/argo-go/kerberos_defaults_windows.go @@ -0,0 +1,19 @@ +//go:build windows + +package main + +import ( + "os" + "path/filepath" +) + +func defaultKerberosConfigPath() string { + if windowsDirectory := os.Getenv("WINDIR"); windowsDirectory != "" { + return filepath.Join(windowsDirectory, "krb5.ini") + } + return `C:\Windows\krb5.ini` +} + +func defaultKerberosCCachePath() string { + return "" +} diff --git a/agents/drivers/argo-go/main.go b/agents/drivers/argo-go/main.go new file mode 100644 index 0000000000..da3216b807 --- /dev/null +++ b/agents/drivers/argo-go/main.go @@ -0,0 +1,620 @@ +package main + +import ( + "bufio" + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "os" + "runtime" + "strconv" + "strings" + "sync" + "time" + _ "time/tzdata" +) + +const ( + protocolVersion = 2 + defaultMaxRows = 10000 + defaultPageSize = 1000 + defaultFetchSize = 50 + legacyAgentSessionID = "__legacy__" + maxAgentSessions = 256 + querySessionIdleTime = 10 * time.Minute +) + +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 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 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 + lastAccessed time.Time +} + +type server struct { + params connectParams + config connectionConfig + + connectionMu sync.Mutex + database *sql.DB + connection *sql.Conn + + querySessions map[string]*querySession + nextSessionID uint64 + + activeMu 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() { + configureRuntimeParallelism() + runtimeServer := newRuntimeServer() + 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() + result, _ := runtimeServer.handleLine(line) + encoderMu.Lock() + _ = encoder.Encode(result) + encoderMu.Unlock() + return + } + requests.Add(1) + go func(line string) { + defer requests.Done() + result, _ := runtimeServer.handleLine(line) + encoderMu.Lock() + defer encoderMu.Unlock() + if err := encoder.Encode(result); err != nil { + fmt.Fprintf(os.Stderr, "failed to write response: %v\n", err) + } + }(line) + } + requests.Wait() +} + +func configureRuntimeParallelism() { + if raw := strings.TrimSpace(os.Getenv("DBX_AGENT_HIVE_GOMAXPROCS")); raw != "" { + if configured, err := strconv.Atoi(raw); err == nil && configured > 0 { + runtime.GOMAXPROCS(configured) + return + } + } + if strings.TrimSpace(os.Getenv("GOMAXPROCS")) == "" { + runtime.GOMAXPROCS(min(runtime.NumCPU(), 4)) + } +} + +func newRuntimeServer() *runtimeServer { + return &runtimeServer{sessions: map[string]*agentSession{}} +} + +func (runtimeServer *runtimeServer) handleLine(line string) (response, bool) { + var request request + if err := json.Unmarshal([]byte(line), &request); err != nil { + return errorResponse(nil, "", "", err), false + } + if len(request.ID) == 0 { + request.ID = json.RawMessage("1") + } + result, shutdown, err := runtimeServer.dispatch(request.Method, request.Params) + if err != nil { + return errorResponse(request.ID, request.Method, stringParam(request.Params, "agentSessionId"), err), false + } + return response{JSONRPC: "2.0", ID: request.ID, Result: result}, shutdown +} + +func (runtimeServer *runtimeServer) dispatch(method string, params map[string]json.RawMessage) (any, bool, error) { + switch method { + case "handshake": + return handshakeResult(true), false, nil + case "open_session": + id := stringParam(params, "agentSessionId") + if id == "" { + return nil, false, errors.New("agentSessionId is required") + } + var connection connectParams + if err := decodeParams(params, &connection); err != nil { + return nil, false, err + } + return map[string]bool{"ok": true}, false, runtimeServer.openSession(id, connection) + case "close_session": + return map[string]bool{"ok": true}, false, runtimeServer.closeSession(stringParam(params, "agentSessionId")) + case "validate_session": + session, err := runtimeServer.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 := runtimeServer.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": + var connection connectParams + if err := decodeParams(params, &connection); err != nil { + return nil, false, err + } + result, err := testConnection(connection) + return result, false, err + case "connect": + var connection connectParams + if err := decodeParams(params, &connection); err != nil { + return nil, false, err + } + _ = runtimeServer.closeSession(legacyAgentSessionID) + return map[string]bool{"ok": true}, false, runtimeServer.openSession(legacyAgentSessionID, connection) + case "disconnect": + return map[string]bool{"ok": true}, false, runtimeServer.closeSession(legacyAgentSessionID) + case "shutdown": + return map[string]bool{"ok": true}, true, runtimeServer.closeAllSessions() + } + + sessionID := stringParam(params, "agentSessionId") + if sessionID == "" { + sessionID = legacyAgentSessionID + } + session, err := runtimeServer.session(sessionID) + if err != nil { + return nil, false, err + } + session.mu.Lock() + defer session.mu.Unlock() + return session.server.dispatch(method, params) +} + +func (runtimeServer *runtimeServer) openSession(id string, params connectParams) error { + runtimeServer.mu.Lock() + if len(runtimeServer.sessions) >= maxAgentSessions { + runtimeServer.mu.Unlock() + return fmt.Errorf("maximum Hive Agent sessions reached (%d)", maxAgentSessions) + } + if _, exists := runtimeServer.sessions[id]; exists { + runtimeServer.mu.Unlock() + return fmt.Errorf("Hive Agent session already exists: %s", id) + } + runtimeServer.mu.Unlock() + + server, err := newServer(params) + if err != nil { + return err + } + runtimeServer.mu.Lock() + defer runtimeServer.mu.Unlock() + if _, exists := runtimeServer.sessions[id]; exists { + _ = server.disconnect() + return fmt.Errorf("Hive Agent session already exists: %s", id) + } + if len(runtimeServer.sessions) >= maxAgentSessions { + _ = server.disconnect() + return fmt.Errorf("maximum Hive Agent sessions reached (%d)", maxAgentSessions) + } + runtimeServer.sessions[id] = &agentSession{server: server} + return nil +} + +func (runtimeServer *runtimeServer) session(id string) (*agentSession, error) { + runtimeServer.mu.RLock() + session := runtimeServer.sessions[id] + runtimeServer.mu.RUnlock() + if session == nil { + return nil, fmt.Errorf("Hive Agent session not found: %s", id) + } + return session, nil +} + +func (runtimeServer *runtimeServer) closeSession(id string) error { + if id == "" { + id = legacyAgentSessionID + } + runtimeServer.mu.Lock() + session := runtimeServer.sessions[id] + delete(runtimeServer.sessions, id) + runtimeServer.mu.Unlock() + if session == nil { + return nil + } + session.mu.Lock() + defer session.mu.Unlock() + return session.server.disconnect() +} + +func (runtimeServer *runtimeServer) closeAllSessions() error { + runtimeServer.mu.Lock() + sessions := runtimeServer.sessions + runtimeServer.sessions = map[string]*agentSession{} + runtimeServer.mu.Unlock() + var failures []string + for id, session := range sessions { + session.mu.Lock() + err := session.server.disconnect() + session.mu.Unlock() + if err != nil { + failures = append(failures, fmt.Sprintf("%s: %v", id, err)) + } + } + if len(failures) > 0 { + return fmt.Errorf("close Hive Agent sessions: %s", strings.Join(failures, "; ")) + } + return nil +} + +func newServer(params connectParams) (*server, error) { + config, err := parseConnectionConfig(params) + if err != nil { + return nil, err + } + server := &server{ + params: params, + config: config, + querySessions: map[string]*querySession{}, + } + if err := server.openConnection(); err != nil { + return nil, err + } + return server, nil +} + +func (server *server) openConnection() error { + server.connectionMu.Lock() + defer server.connectionMu.Unlock() + if server.connection != nil { + return nil + } + database := openHiveDatabase(server.config) + ctx, cancel := context.WithTimeout(context.Background(), server.connectionOpenTimeout()) + var connection *sql.Conn + connection, err := database.Conn(ctx) + if err == nil { + err = connection.PingContext(ctx) + } + cancel() + if err == nil { + err = runHiveInitStatements(context.Background(), connection, server.config.InitStatements, server.config.FetchSize) + } + if err != nil { + if connection != nil { + _ = connection.Close() + } + _ = database.Close() + return err + } + server.database = database + server.connection = connection + return nil +} + +func (server *server) connectionOpenTimeout() time.Duration { + timeout := server.config.ConnectTimeout + if strings.EqualFold(server.config.Auth, "BROWSER") && strings.TrimSpace(server.config.BrowserToken) == "" { + timeout += server.config.BrowserResponseTimeout + } + return timeout +} + +func runHiveInitStatements(ctx context.Context, connection *sql.Conn, statements []string, fetchSize int) error { + for _, statement := range statements { + rows, _, hasResultSet, err := executeHiveStatement(ctx, connection, statement, fetchSize) + if err != nil { + return fmt.Errorf("execute Hive initFile statement: %w", err) + } + if !hasResultSet { + continue + } + for rows.Next() { + } + if rowsErr := rows.Err(); rowsErr != nil { + _ = rows.Close() + return fmt.Errorf("read Hive initFile statement result: %w", rowsErr) + } + if closeErr := rows.Close(); closeErr != nil { + return fmt.Errorf("close Hive initFile statement result: %w", closeErr) + } + } + return nil +} + +func (server *server) dispatch(method string, params map[string]json.RawMessage) (any, bool, error) { + switch method { + case "validate_connection": + return map[string]bool{"ok": true}, false, server.validateConnection() + case "connection_info": + result, err := server.connectionInfo() + return result, false, err + case "list_databases": + result, err := server.listDatabases() + return result, false, err + case "list_schemas": + result, err := server.listSchemas(stringSliceParam(params, "visible_schemas")) + return result, false, err + case "list_tables": + result, err := server.listTables(stringParam(params, "schema"), metadataListConstraintsFromParams(params)) + return result, false, err + case "get_table_comment": + result, err := server.getTableComment(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "list_objects": + result, err := server.listObjects( + stringParam(params, "database"), + stringParam(params, "schema"), + metadataListConstraintsFromParams(params), + ) + return result, false, err + case "list_data_types": + result, err := server.listDataTypes() + return result, false, err + case "completion_assistant_search_v1": + var input completionAssistantRequest + if err := decodeParams(params, &input); err != nil { + return nil, false, err + } + result, err := server.completionAssistantSearch(input) + return result, false, err + case "get_columns": + result, err := server.getColumns(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "list_indexes": + return []indexInfo{}, false, nil + case "list_foreign_keys": + return []foreignKeyInfo{}, false, nil + case "list_triggers": + return []triggerInfo{}, false, nil + case "list_constraints", "list_partitions", "list_subpartitions": + return []any{}, false, nil + case "get_object_source": + result, err := server.getObjectSource( + stringParam(params, "database"), + stringParam(params, "schema"), + firstNonEmpty(stringParam(params, "name"), stringParam(params, "table")), + stringParam(params, "object_type"), + ) + return result, false, err + case "get_table_ddl": + result, err := server.getTableDDL(stringParam(params, "schema"), stringParam(params, "table")) + return result, false, err + case "get_explain_info": + result, err := server.getExplainInfo(stringParam(params, "sql")) + return map[string]any{"plan": result, "has_actual_stats": false}, false, err + case "execute_query": + result, err := server.executeQuery(queryOptionsFromParams(params)) + return result, false, err + case "execute_query_page", "start_table_read": + result, err := server.executeQueryPage(queryOptionsFromParams(params), intParam(params, "pageSize")) + return result, false, err + case "fetch_query_page", "fetch_table_read_page": + result, err := server.fetchQueryPage(stringParam(params, "sessionId"), intParam(params, "pageSize")) + return result, false, err + case "close_query_session", "close_table_read_session": + return server.closeQuerySession(stringParam(params, "sessionId")), false, nil + case "execute_transaction": + result, err := server.executeStatements(params, true) + return result, false, err + case "execute_batch": + result, err := server.executeStatements(params, false) + return result, false, err + case "disconnect": + return map[string]bool{"ok": true}, false, server.disconnect() + case "shutdown": + return map[string]bool{"ok": true}, true, server.disconnect() + default: + return nil, false, fmt.Errorf("unknown method: %s", method) + } +} + +func handshakeResult(multiSession bool) map[string]any { + capabilities := []string{ + "connect", "test_connection", "metadata", "query", "paged_query", "transaction", "ddl", "structured_error_v1", + } + if multiSession { + capabilities = append(capabilities, "multi_session") + } + return map[string]any{ + "protocolVersion": protocolVersion, + "agentProtocolVersion": protocolVersion, + "capabilities": capabilities, + } +} + +func testConnection(params connectParams) (map[string]any, error) { + server, err := newServer(params) + if err != nil { + return nil, err + } + defer server.disconnect() + if err := server.validateConnection(); err != nil { + return nil, err + } + info, err := server.connectionInfo() + if err != nil { + return nil, err + } + return map[string]any{"ok": true, "info": info}, nil +} + +func (server *server) disconnect() error { + server.cancelActiveQuery() + if err := server.closeAllQuerySessions(); err != nil { + return err + } + server.connectionMu.Lock() + connection := server.connection + database := server.database + server.connection = nil + server.database = nil + server.connectionMu.Unlock() + var failures []string + if connection != nil { + if err := connection.Close(); err != nil { + failures = append(failures, err.Error()) + } + } + if database != nil { + if err := database.Close(); err != nil { + failures = append(failures, err.Error()) + } + } + if len(failures) > 0 { + return errors.New(strings.Join(failures, "; ")) + } + return nil +} + +func (server *server) requireConnection() (*sql.Conn, error) { + server.connectionMu.Lock() + connection := server.connection + server.connectionMu.Unlock() + if connection == nil { + return nil, errors.New("Hive connection is not open") + } + return connection, nil +} + +func (server *server) setActiveOperation(cancel context.CancelFunc) { + server.activeMu.Lock() + server.activeCancel = cancel + server.activeMu.Unlock() +} + +func (server *server) clearActiveOperation(cancel context.CancelFunc) { + cancel() + server.activeMu.Lock() + server.activeCancel = nil + server.activeMu.Unlock() +} + +func (server *server) cancelActiveQuery() { + server.activeMu.Lock() + cancel := server.activeCancel + server.activeMu.Unlock() + if cancel != nil { + cancel() + } +} + +func queryOptionsFromParams(params map[string]json.RawMessage) queryOptions { + return queryOptions{ + SQL: stringParam(params, "sql"), + Database: stringParam(params, "database"), + Schema: stringParam(params, "schema"), + MaxRows: intParam(params, "maxRows"), + FetchSize: intParam(params, "fetchSize"), + TimeoutSecs: intParam(params, "timeoutSecs"), + } +} + +func 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 { + if raw, ok := params[key]; ok { + var value string + if json.Unmarshal(raw, &value) == nil { + return value + } + } + return "" +} + +func intParam(params map[string]json.RawMessage, key string) int { + if raw, ok := params[key]; ok { + var value int + if json.Unmarshal(raw, &value) == nil { + return value + } + } + return 0 +} + +func stringSliceParam(params map[string]json.RawMessage, key string) []string { + raw, ok := params[key] + if !ok { + return nil + } + var value []string + if json.Unmarshal(raw, &value) == nil { + return value + } + return nil +} + +func errorResponse(id json.RawMessage, method, sessionID string, err error) response { + return response{JSONRPC: "2.0", ID: id, Error: classifyRPCError(method, sessionID, err)} +} diff --git a/agents/drivers/argo-go/metadata.go b/agents/drivers/argo-go/metadata.go new file mode 100644 index 0000000000..c59c149d11 --- /dev/null +++ b/agents/drivers/argo-go/metadata.go @@ -0,0 +1,1073 @@ +package main + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log" + "sort" + "strconv" + "strings" + + "github.com/t8y2/dbx/agents/go-common/gohive" +) + +const metadataQueryLimit = 100000 + +var hiveTypes = []string{ + "tinyint", "smallint", "int", "bigint", "boolean", "float", "double", "decimal", "string", "varchar", + "char", "binary", "date", "timestamp", "timestamp with local time zone", "interval_year_month", + "interval_day_time", "array", "map", "struct", "uniontype", "void", +} + +type databaseInfo struct { + Name string `json:"name"` +} + +type tableInfo struct { + Name string `json:"name"` + TableType string `json:"table_type"` + Comment *string `json:"comment"` + ParentSchema *string `json:"parent_schema,omitempty"` + ParentName *string `json:"parent_name,omitempty"` +} + +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 objectSource struct { + Name string `json:"name"` + ObjectType string `json:"object_type"` + Schema *string `json:"schema"` + Source string `json:"source"` +} + +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 (value indexInfo) MarshalJSON() ([]byte, error) { + type alias indexInfo + copy := alias(value) + if copy.Columns == nil { + copy.Columns = []string{} + } + if copy.IncludedColumns == nil { + copy.IncludedColumns = []string{} + } + return json.Marshal(copy) +} + +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"` +} + +type metadataListConstraints struct { + Filter string + Limit int + Offset int + ObjectTypes []string +} + +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"` +} + +func hiveDataTypes() []string { + return append([]string(nil), hiveTypes...) +} + +type hiveMetadataRows struct { + indexes map[string]int + rows [][]any +} + +func newHiveMetadataRows(result gohive.MetadataResult) hiveMetadataRows { + indexes := make(map[string]int, len(result.Columns)) + for index, column := range result.Columns { + indexes[normalizeMetadataColumn(column)] = index + } + rows := make([][]any, 0, len(result.Rows)) + for _, row := range result.Rows { + values := make([]any, len(row)) + for index, value := range row { + values[index] = value + } + rows = append(rows, values) + } + return hiveMetadataRows{indexes: indexes, rows: rows} +} + +func (rows hiveMetadataRows) value(row []any, names ...string) any { + for _, name := range names { + if index, ok := rows.indexes[normalizeMetadataColumn(name)]; ok && index >= 0 && index < len(row) { + return row[index] + } + } + return nil +} + +func normalizeMetadataColumn(value string) string { + return strings.NewReplacer("_", "", "-", "", " ", "").Replace(strings.ToUpper(strings.TrimSpace(value))) +} + +func (server *server) hiveMetadata(operation func(context.Context, gohive.MetadataProvider) (gohive.MetadataResult, error)) (gohive.MetadataResult, error) { + connection, err := server.requireConnection() + if err != nil { + return gohive.MetadataResult{}, err + } + ctx, cancel := context.WithCancel(context.Background()) + server.setActiveOperation(cancel) + defer server.clearActiveOperation(cancel) + var result gohive.MetadataResult + err = connection.Raw(func(rawConnection any) error { + provider, ok := rawConnection.(gohive.MetadataProvider) + if !ok { + return errors.New("Hive driver does not expose HiveServer2 metadata RPCs") + } + var operationErr error + result, operationErr = operation(ctx, provider) + return operationErr + }) + return result, err +} + +func (server *server) connectionInfo() (map[string]any, error) { + version := "" + if result, err := server.executeQuery(queryOptions{SQL: "SELECT VERSION()", MaxRows: 1, TimeoutSecs: 5}); err == nil && len(result.Rows) > 0 && len(result.Rows[0]) > 0 { + version = stringValue(result.Rows[0][0]) + } + username := server.config.Username + if result, err := server.executeQuery(queryOptions{SQL: "SELECT CURRENT_USER()", MaxRows: 1, TimeoutSecs: 5}); err == nil && len(result.Rows) > 0 && len(result.Rows[0]) > 0 { + if current := stringValue(result.Rows[0][0]); current != "" { + username = current + } + } + // The argo agent exclusively serves 星环Argo (Transwarp ArgoDB) connections; vanilla + // Hive/Kyuubi/Impala stay on hive-go. Brand accordingly so the connection info + // panel reflects the actual server family. + productName := "ArgoDB (Transwarp)" + compatibilityMode := "argo" + driverName := "DBX ArgoDB Go Agent" + return map[string]any{ + "database": server.config.Database, + "schema": server.config.Database, + "username": username, + "version": version, + "sqlDialect": "HIVE", + "identifierQuote": "`", + "compatibilityMode": compatibilityMode, + "databaseInfo": map[string]string{ + "productName": productName, + "productVersion": version, + "unquotedIdentifierCase": "mixed", + "quotedIdentifierCase": "mixed", + "driverName": driverName, + "driverVersion": "gohive-v2.1.0", + }, + }, nil +} + +func (server *server) listDatabases() ([]databaseInfo, error) { + result, err := server.executeQuery(queryOptions{SQL: "SHOW DATABASES", MaxRows: metadataQueryLimit}) + if err == nil { + return databaseInfoFromQueryRows(result.Rows), nil + } + metadataResult, metadataErr := server.hiveMetadata(func(ctx context.Context, provider gohive.MetadataProvider) (gohive.MetadataResult, error) { + return provider.GetHiveSchemas(ctx, "%") + }) + if metadataErr != nil { + return nil, fmt.Errorf("SHOW DATABASES failed (%v); HiveServer2 metadata fallback failed: %w", err, metadataErr) + } + rows := newHiveMetadataRows(metadataResult) + values := make([]databaseInfo, 0, len(rows.rows)) + seen := map[string]bool{} + for _, row := range rows.rows { + name := metadataString(rows.value(row, "TABLE_SCHEM", "SCHEMA_NAME")) + if name == "" || seen[name] { + continue + } + seen[name] = true + values = append(values, databaseInfo{Name: name}) + } + sort.Slice(values, func(first, second int) bool { return values[first].Name < values[second].Name }) + return values, nil +} + +func databaseInfoFromQueryRows(rows [][]any) []databaseInfo { + values := make([]databaseInfo, 0, len(rows)) + seen := map[string]bool{} + for _, row := range rows { + name := firstRowValue(row) + if name == "" || seen[name] { + continue + } + seen[name] = true + values = append(values, databaseInfo{Name: name}) + } + sort.Slice(values, func(first, second int) bool { return values[first].Name < values[second].Name }) + return values +} + +func (server *server) listSchemas(visibleSchemas []string) ([]string, error) { + if visibleSchemas != nil && len(visibleSchemas) == 0 { + return []string{}, nil + } + databases, err := server.listDatabases() + if err != nil { + return nil, err + } + visible := map[string]bool{} + for _, schema := range visibleSchemas { + visible[schema] = true + } + values := make([]string, 0, len(databases)) + for _, database := range databases { + if visibleSchemas != nil && !visible[database.Name] { + continue + } + values = append(values, database.Name) + } + return values, nil +} + +func (server *server) listTables(schema string, constraints metadataListConstraints) ([]tableInfo, error) { + schema = firstNonEmpty(schema, server.config.Database) + requestedTypes := hiveTableTypes(constraints.ObjectTypes) + if len(constraints.ObjectTypes) > 0 && len(requestedTypes) == 0 { + return []tableInfo{}, nil + } + metadataResult, metadataErr := server.hiveMetadata(func(ctx context.Context, provider gohive.MetadataProvider) (gohive.MetadataResult, error) { + return provider.GetHiveTables(ctx, schema, "%", requestedTypes) + }) + if metadataErr == nil { + rows := newHiveMetadataRows(metadataResult) + values := make([]tableInfo, 0, len(rows.rows)) + for _, row := range rows.rows { + name := metadataString(rows.value(row, "TABLE_NAME")) + if name == "" || !metadataNameMatches(name, constraints.Filter) { + continue + } + values = append(values, tableInfo{ + Name: name, + TableType: normalizeHiveTableType(metadataString(rows.value(row, "TABLE_TYPE"))), + Comment: optionalString(metadataString(rows.value(row, "REMARKS", "COMMENT"))), + }) + } + sort.Slice(values, func(first, second int) bool { return values[first].Name < values[second].Name }) + return applyMetadataWindow(values, constraints.Offset, constraints.Limit), nil + } + type fallbackQuery struct { + operation string + statement string + objectType string + } + fallbackQueries := make([]fallbackQuery, 0, 2) + if containsString(requestedTypes, "TABLE") { + fallbackQueries = append(fallbackQueries, fallbackQuery{ + operation: "SHOW TABLES", + statement: "SHOW TABLES IN " + quoteHiveIdentifier(schema), + objectType: "TABLE", + }) + } + if containsString(requestedTypes, "VIEW") || containsString(requestedTypes, "MATERIALIZED VIEW") { + fallbackQueries = append(fallbackQueries, fallbackQuery{ + operation: "SHOW VIEWS", + statement: "SHOW VIEWS IN " + quoteHiveIdentifier(schema), + objectType: "VIEW", + }) + } + objectsByName := make(map[string]tableInfo) + tableFallbackSucceeded := false + for _, fallback := range fallbackQueries { + result, err := server.executeQuery(queryOptions{SQL: fallback.statement, MaxRows: metadataQueryLimit}) + if err != nil { + // Older Hive and Impala versions can list tables but do not support SHOW VIEWS. + // Keep the usable table result for mixed requests; explicit view requests still fail. + if fallback.objectType == "VIEW" && tableFallbackSucceeded && showViewsUnsupported(err) { + continue + } + return nil, fmt.Errorf("HiveServer2 metadata failed (%v); %s fallback failed: %w", metadataErr, fallback.operation, err) + } + if fallback.objectType == "TABLE" { + tableFallbackSucceeded = true + } + for _, row := range result.Rows { + name := showTablesRowName(result.Columns, row) + if name == "" || !metadataNameMatches(name, constraints.Filter) { + continue + } + candidate := tableInfo{Name: name, TableType: fallback.objectType, Comment: nil} + if existing, ok := objectsByName[name]; ok && existing.TableType == "VIEW" && candidate.TableType != "VIEW" { + continue + } + objectsByName[name] = candidate + } + } + values := make([]tableInfo, 0, len(objectsByName)) + for _, value := range objectsByName { + values = append(values, value) + } + sort.Slice(values, func(first, second int) bool { return values[first].Name < values[second].Name }) + return applyMetadataWindow(values, constraints.Offset, constraints.Limit), nil +} + +func showViewsUnsupported(err error) bool { + if err == nil { + return false + } + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return false + } + message := strings.ToLower(err.Error()) + operationalMarkers := []string{ + "permission", + "access denied", + "not authorized", + "unauthorized", + "authentication", + "authorization", + "timeout", + "timed out", + "deadline exceeded", + "cancelled", + "canceled", + "transport", + "connection", + "broken pipe", + "network", + } + for _, marker := range operationalMarkers { + if strings.Contains(message, marker) { + return false + } + } + explicitMarkers := []string{ + "unsupported", + "not supported", + "not implemented", + "unknown statement", + "unrecognized statement", + } + for _, marker := range explicitMarkers { + if strings.Contains(message, marker) { + return true + } + } + parseMarkers := []string{ + "parseexception", + "parse error", + "syntax error", + "mismatched input", + "cannot recognize input", + "no viable alternative", + } + for _, marker := range parseMarkers { + if strings.Contains(message, marker) { + return true + } + } + return false +} + +func (server *server) listObjects(database, schema string, constraints metadataListConstraints) ([]objectInfo, error) { + if !acceptsHiveTable(constraints.ObjectTypes) && !(server.supportsRoutines() && acceptsHiveRoutine(constraints.ObjectTypes)) { + return []objectInfo{}, nil + } + schema = firstNonEmpty(schema, server.config.Database) + values := make([]objectInfo, 0) + if acceptsHiveTable(constraints.ObjectTypes) { + tables, err := server.listTables(schema, constraints) + if err != nil { + return nil, err + } + for _, table := range tables { + values = append(values, objectInfo{Name: table.Name, ObjectType: table.TableType, Schema: schema, Comment: table.Comment}) + } + } + if server.supportsRoutines() && acceptsRoutineType(constraints.ObjectTypes, "PROCEDURE") { + procedures, err := server.listRoutines(database, schema, constraints, "PROCEDURE") + if err != nil { + return nil, err + } + values = append(values, procedures...) + } + if server.supportsRoutines() && acceptsRoutineType(constraints.ObjectTypes, "FUNCTION") { + functions, err := server.listRoutines(database, schema, constraints, "FUNCTION") + if err != nil { + return nil, err + } + values = append(values, functions...) + } + sort.Slice(values, func(first, second int) bool { return values[first].Name < values[second].Name }) + return applyMetadataWindow(values, constraints.Offset, constraints.Limit), nil +} + +// supportsRoutines reports whether the connected server is known to expose the +// Hive procedure / function catalog views (system.procedures_v / +// system.functions_v). This agent exclusively serves 星环Argo (Transwarp ArgoDB), +// which ships those views, so routine support is unconditional — unlike hive-go, +// which still gates it on the connection's database_type for vanilla Hive. +func (server *server) supportsRoutines() bool { + return true +} + +// routineDatabaseCandidates returns the database_name filters to try against +// system.procedures_v / system.functions_v. The sidebar may pass the active +// database through either the schema or database RPC parameter, so mirror the +// JDBC plugin and try both. The connection default is only a candidate when +// the caller supplied neither parameter: appending it unconditionally would +// return the default database's routines under an explicit database node that +// has none of its own. +func routineDatabaseCandidates(schema, database, connectionDatabase string) []string { + seen := map[string]bool{} + values := make([]string, 0, 3) + add := func(candidate string) { + candidate = strings.TrimSpace(candidate) + if candidate == "" { + return + } + key := strings.ToLower(candidate) + if seen[key] { + return + } + seen[key] = true + values = append(values, candidate) + } + add(database) + add(schema) + if strings.TrimSpace(database) == "" && strings.TrimSpace(schema) == "" { + add(connectionDatabase) + } + return values +} + +// listRoutines queries the server's procedure / function catalog views when the +// driver supports them. Hive and most forks (ArgoDB, Inceptor, Transwarp) expose +// stored procedures / functions through the system.procedures_v / system.functions_v +// views, with columns (procedure_name | function_name, database_name, full_text, ...). +// The full_text column carries the routine source used by getObjectSource. +// +// The query is best-effort: when the view is missing or the server rejects it +// (older Hive without procedure support), the call returns an empty slice and +// nil error so the caller can fall back to listing tables. +func (server *server) listRoutines(database, schema string, constraints metadataListConstraints, routineType string) ([]objectInfo, error) { + nameColumn := "procedure_name" + viewName := "system.procedures_v" + if strings.EqualFold(routineType, "FUNCTION") { + nameColumn = "function_name" + viewName = "system.functions_v" + } + candidates := routineDatabaseCandidates(schema, database, server.config.Database) + if len(candidates) == 0 { + return []objectInfo{}, nil + } + likePattern := buildRoutineLikePattern(constraints.Filter) + resolvedSchema := firstNonEmpty(schema, database, server.config.Database) + for _, targetSchema := range candidates { + // database_name is a string column, so the schema filter must be a single-quoted + // literal — not a backtick-quoted identifier. ArgoDB/Inceptor reject lower(`ods`) + // against system.procedures_v with an ERROR_STATUS, while lower('ods') works. + schemaLiteral := "'" + strings.ReplaceAll(targetSchema, "'", "''") + "'" + sql := "SELECT " + nameColumn + " FROM " + viewName + + " WHERE lower(database_name) = lower(" + schemaLiteral + ")" + + " AND lower(" + nameColumn + ") LIKE " + likePattern + + " ORDER BY " + nameColumn + result, err := server.executeQuery(queryOptions{SQL: sql, MaxRows: metadataQueryLimit}) + if err != nil { + log.Printf( + "[argo-go][listRoutines] query failed: database=%q schema=%q routineType=%s sql=%q err=%v", + targetSchema, + resolvedSchema, + routineType, + sql, + err, + ) + continue + } + values := make([]objectInfo, 0, len(result.Rows)) + for _, row := range result.Rows { + name := rowString(row, 0) + if name == "" { + continue + } + values = append(values, objectInfo{Name: name, ObjectType: strings.ToUpper(routineType), Schema: resolvedSchema, Comment: nil}) + } + if len(values) > 0 { + return values, nil + } + } + return []objectInfo{}, nil +} + +// buildRoutineLikePattern turns a user-supplied filter into a Hive-safe LIKE +// literal: wraps with %, escapes \, %, _ (the LIKE metacharacters), and quotes +// the whole literal so it can be concatenated directly into SQL. +func buildRoutineLikePattern(filter string) string { + escaped := strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`).Replace(filter) + return "'%" + escaped + "%'" +} + +// acceptsHiveRoutine reports whether the requested object types include any +// routine (PROCEDURE / FUNCTION) that hive-go needs to surface through listObjects. +func acceptsHiveRoutine(objectTypes []string) bool { + for _, objectType := range objectTypes { + if strings.EqualFold(objectType, "PROCEDURE") || strings.EqualFold(objectType, "FUNCTION") { + return true + } + } + return false +} + +// acceptsRoutineType reports whether the requested object types include the +// given routine kind (case-insensitive). +func acceptsRoutineType(objectTypes []string, routineType string) bool { + for _, objectType := range objectTypes { + if strings.EqualFold(objectType, routineType) { + return true + } + } + return false +} + +func (server *server) getColumns(schema, table string) ([]columnInfo, error) { + if strings.TrimSpace(table) == "" { + return nil, errors.New("table is required") + } + schema = firstNonEmpty(schema, server.config.Database) + metadataResult, metadataErr := server.hiveMetadata(func(ctx context.Context, provider gohive.MetadataProvider) (gohive.MetadataResult, error) { + return provider.GetHiveColumns(ctx, schema, table, "%") + }) + if metadataErr == nil { + rows := newHiveMetadataRows(metadataResult) + values := make([]columnInfo, 0, len(rows.rows)) + for _, row := range rows.rows { + name := metadataString(rows.value(row, "COLUMN_NAME")) + if name == "" { + continue + } + dataType := metadataString(rows.value(row, "TYPE_NAME")) + columnSize := metadataIntPointer(rows.value(row, "COLUMN_SIZE")) + values = append(values, columnInfo{ + Name: name, + DataType: dataType, + IsNullable: metadataNullable(rows.value(row, "NULLABLE", "IS_NULLABLE")), + ColumnDefault: optionalString(metadataString(rows.value(row, "COLUMN_DEF"))), + Comment: optionalString(metadataString(rows.value(row, "REMARKS", "COMMENT"))), + NumericPrecision: columnSize, + NumericScale: metadataIntPointer(rows.value(row, "DECIMAL_DIGITS")), + CharacterMaximumLength: characterLengthForType(dataType, columnSize), + }) + } + return values, nil + } + qualified := qualifiedHiveName(schema, table) + result, err := server.executeQuery(queryOptions{SQL: "DESCRIBE " + qualified, MaxRows: metadataQueryLimit}) + if err != nil { + return nil, fmt.Errorf("HiveServer2 metadata failed (%v); DESCRIBE fallback failed: %w", metadataErr, err) + } + values := make([]columnInfo, 0, len(result.Rows)) + for _, row := range result.Rows { + name := rowString(row, 0) + if name == "" || strings.HasPrefix(name, "#") { + continue + } + dataType := rowString(row, 1) + comment := optionalString(rowString(row, 2)) + values = append(values, columnInfo{ + Name: name, + DataType: dataType, + IsNullable: true, + Comment: comment, + }) + } + return values, nil +} + +func (server *server) getTableComment(schema, table string) (*string, error) { + if strings.TrimSpace(table) == "" { + return nil, errors.New("table is required") + } + schema = firstNonEmpty(schema, server.config.Database) + metadataResult, err := server.hiveMetadata(func(ctx context.Context, provider gohive.MetadataProvider) (gohive.MetadataResult, error) { + return provider.GetHiveTables(ctx, schema, table, nil) + }) + if err != nil { + tables, listErr := server.listTables(schema, metadataListConstraints{Filter: table}) + if listErr != nil { + return nil, fmt.Errorf("HiveServer2 table comment metadata failed (%v); table listing fallback failed: %w", err, listErr) + } + for _, candidate := range tables { + if strings.EqualFold(candidate.Name, table) { + return candidate.Comment, nil + } + } + return nil, nil + } + rows := newHiveMetadataRows(metadataResult) + for _, row := range rows.rows { + if strings.EqualFold(metadataString(rows.value(row, "TABLE_NAME")), table) { + return optionalString(metadataString(rows.value(row, "REMARKS", "COMMENT"))), nil + } + } + return nil, nil +} + +func (server *server) listDataTypes() ([]string, error) { + metadataResult, err := server.hiveMetadata(func(ctx context.Context, provider gohive.MetadataProvider) (gohive.MetadataResult, error) { + return provider.GetHiveTypeInfo(ctx) + }) + if err != nil { + return hiveDataTypes(), nil + } + rows := newHiveMetadataRows(metadataResult) + values := make([]string, 0, len(rows.rows)) + seen := map[string]bool{} + for _, row := range rows.rows { + name := strings.ToLower(metadataString(rows.value(row, "TYPE_NAME"))) + if name == "" || seen[name] { + continue + } + seen[name] = true + values = append(values, name) + } + if len(values) == 0 { + return hiveDataTypes(), nil + } + sort.Strings(values) + return values, nil +} + +func (server *server) getTableDDL(schema, table string) (string, error) { + if strings.TrimSpace(table) == "" { + return "", errors.New("table is required") + } + result, err := server.executeQuery(queryOptions{ + SQL: "SHOW CREATE TABLE " + qualifiedHiveName(schema, table), + MaxRows: metadataQueryLimit, + }) + if err != nil { + return "", err + } + lines := make([]string, 0, len(result.Rows)) + for _, row := range result.Rows { + if line := firstRowValue(row); line != "" { + lines = append(lines, line) + } + } + if len(lines) == 0 { + return "", nil + } + return strings.Join(lines, "\n") + "\n", nil +} + +func (server *server) getObjectSource(database, schema, name, objectType string) (objectSource, error) { + schema = firstNonEmpty(schema, database, server.config.Database) + var source string + var err error + switch strings.ToUpper(objectType) { + case "PROCEDURE", "FUNCTION": + if !server.supportsRoutines() { + return objectSource{}, fmt.Errorf("routine source is not supported for %s connections", server.params.DatabaseType) + } + source, err = server.getRoutineSource(database, schema, name, strings.ToUpper(objectType)) + default: + source, err = server.getTableDDL(schema, name) + } + if err != nil { + return objectSource{}, err + } + return objectSource{ + Name: name, + ObjectType: strings.ToUpper(objectType), + Schema: optionalString(schema), + Source: source, + }, nil +} + +// getRoutineSource fetches a procedure or function's full source from the +// server's system.procedures_v / system.functions_v view. Returns an empty +// string when the view is missing or the routine is not found, so callers can +// fall back to other sources. full_text may span multiple rows (the underlying +// query driver splits long strings), so we join them like getTableDDL does. +func (server *server) getRoutineSource(database, schema, name, routineType string) (string, error) { + nameColumn := "procedure_name" + viewName := "system.procedures_v" + if strings.EqualFold(routineType, "FUNCTION") { + nameColumn = "function_name" + viewName = "system.functions_v" + } + candidates := routineDatabaseCandidates(schema, database, server.config.Database) + if len(candidates) == 0 { + return "", nil + } + escapedName := strings.ReplaceAll(name, "'", "''") + for _, targetSchema := range candidates { + sql := "SELECT full_text FROM " + viewName + + " WHERE lower(database_name) = lower('" + strings.ReplaceAll(targetSchema, "'", "''") + "')" + + " AND " + nameColumn + " = '" + escapedName + "'" + result, err := server.executeQuery(queryOptions{SQL: sql, MaxRows: metadataQueryLimit}) + if err != nil { + log.Printf( + "[argo-go][getRoutineSource] query failed: database=%q schema=%q name=%q routineType=%s sql=%q err=%v", + targetSchema, + schema, + name, + routineType, + sql, + err, + ) + continue + } + lines := make([]string, 0, len(result.Rows)) + for _, row := range result.Rows { + if line := firstRowValue(row); line != "" { + lines = append(lines, line) + } + } + if len(lines) == 0 { + continue + } + return strings.Join(lines, "\n") + "\n", nil + } + return "", nil +} + +func (server *server) getExplainInfo(sqlText string) (string, error) { + sqlText = trimStatementSQL(sqlText) + if sqlText == "" { + return "", errors.New("SQL is required") + } + result, err := server.executeQuery(queryOptions{SQL: "EXPLAIN " + sqlText, MaxRows: metadataQueryLimit}) + if err != nil { + return "", err + } + lines := make([]string, 0, len(result.Rows)) + for _, row := range result.Rows { + lines = append(lines, firstRowValue(row)) + } + return strings.Join(lines, "\n"), nil +} + +func (server *server) completionAssistantSearch(input completionAssistantRequest) (completionAssistantResponse, error) { + maxResults := input.MaxResults + if maxResults <= 0 { + maxResults = 200 + } + schemas := []string{firstNonEmpty(input.Schema, input.Database, server.config.Database)} + if input.GlobalSearch { + listed, err := server.listSchemas(nil) + if err != nil { + return completionAssistantResponse{}, err + } + schemas = listed + } + values := make([]completionAssistantCandidate, 0, maxResults) + incomplete := false + for _, schema := range schemas { + tables, err := server.listTables(schema, metadataListConstraints{Limit: maxResults}) + if err != nil { + return completionAssistantResponse{}, err + } + for _, table := range tables { + if !completionNameMatches(table.Name, input) { + continue + } + schemaCopy := schema + values = append(values, completionAssistantCandidate{ + Name: table.Name, Kind: "table", Database: &schemaCopy, Schema: &schemaCopy, Comment: table.Comment, + }) + if len(values) >= maxResults { + incomplete = true + break + } + } + if len(values) >= maxResults { + break + } + } + return completionAssistantResponse{Candidates: values, Incomplete: incomplete, FallbackUsed: false}, 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, "objectTypes"), + } +} + +func qualifiedHiveName(schema, table string) string { + if strings.TrimSpace(schema) == "" { + return quoteHiveIdentifier(table) + } + return quoteHiveIdentifier(schema) + "." + quoteHiveIdentifier(table) +} + +func metadataNameMatches(name, filter string) bool { + return filter == "" || strings.Contains(strings.ToLower(name), strings.ToLower(filter)) +} + +func acceptsHiveTable(objectTypes []string) bool { + if len(objectTypes) == 0 { + return true + } + for _, objectType := range objectTypes { + if strings.EqualFold(objectType, "table") || strings.EqualFold(objectType, "view") || strings.EqualFold(objectType, "materialized view") { + return true + } + } + return false +} + +func containsString(values []string, expected string) bool { + for _, value := range values { + if value == expected { + return true + } + } + return false +} + +func hiveTableTypes(objectTypes []string) []string { + if len(objectTypes) == 0 { + return []string{"TABLE", "VIEW", "MATERIALIZED VIEW"} + } + values := make([]string, 0, len(objectTypes)) + seen := map[string]bool{} + for _, objectType := range objectTypes { + normalized := strings.ToUpper(strings.TrimSpace(objectType)) + switch normalized { + case "TABLE", "EXTERNAL TABLE", "MANAGED TABLE": + normalized = "TABLE" + case "VIEW": + normalized = "VIEW" + case "MATERIALIZED VIEW", "MATERIALIZED_VIEW": + normalized = "MATERIALIZED VIEW" + default: + continue + } + if !seen[normalized] { + seen[normalized] = true + values = append(values, normalized) + } + } + return values +} + +func normalizeHiveTableType(value string) string { + if strings.Contains(strings.ToUpper(value), "VIEW") { + return "VIEW" + } + return "TABLE" +} + +func metadataString(value any) string { + switch typed := value.(type) { + case nil: + return "" + case string: + return strings.TrimSpace(typed) + case []byte: + return strings.TrimSpace(string(typed)) + default: + return strings.TrimSpace(fmt.Sprint(typed)) + } +} + +func metadataIntPointer(value any) *int { + var parsed int64 + switch typed := value.(type) { + case nil: + return nil + case int: + parsed = int64(typed) + case int8: + parsed = int64(typed) + case int16: + parsed = int64(typed) + case int32: + parsed = int64(typed) + case int64: + parsed = typed + case float32: + parsed = int64(typed) + case float64: + parsed = int64(typed) + default: + value, err := strconv.ParseInt(metadataString(value), 10, 64) + if err != nil { + return nil + } + parsed = value + } + if parsed < 0 || parsed > int64(^uint(0)>>1) { + return nil + } + converted := int(parsed) + return &converted +} + +func metadataNullable(value any) bool { + if parsed := metadataIntPointer(value); parsed != nil { + return *parsed != 0 + } + switch strings.ToUpper(metadataString(value)) { + case "NO", "FALSE", "NOT NULL": + return false + default: + return true + } +} + +func characterLengthForType(dataType string, size *int) *int { + normalized := strings.ToLower(dataType) + if strings.Contains(normalized, "char") || strings.Contains(normalized, "text") || strings.Contains(normalized, "string") { + return size + } + return nil +} + +func applyMetadataWindow[T any](values []T, offset, limit int) []T { + if offset < 0 { + offset = 0 + } + if offset >= len(values) { + return []T{} + } + end := len(values) + if limit > 0 && offset+limit < end { + end = offset + limit + } + return values[offset:end] +} + +func completionNameMatches(name string, input completionAssistantRequest) bool { + mask := input.Mask + if mask == "" { + return true + } + if !input.CaseSensitive { + name = strings.ToLower(name) + mask = strings.ToLower(mask) + } + switch strings.ToLower(input.MatchMode) { + case "exact": + return name == mask + case "prefix": + return strings.HasPrefix(name, mask) + default: + return strings.Contains(name, mask) + } +} + +func firstRowValue(row []any) string { + for _, value := range row { + if text := stringValue(value); text != "" { + return text + } + } + return "" +} + +func showTablesRowName(columns []string, row []any) string { + for index, column := range columns { + normalized := strings.NewReplacer("_", "", "-", "", " ", "").Replace(strings.ToLower(column)) + if normalized == "tablename" || normalized == "tabname" { + if value := rowString(row, index); value != "" { + return value + } + } + } + if len(row) > 1 { + if value := rowString(row, 1); value != "" { + return value + } + } + return firstRowValue(row) +} + +func rowString(row []any, index int) string { + if index < 0 || index >= len(row) { + return "" + } + return strings.TrimSpace(stringValue(row[index])) +} + +func stringValue(value any) string { + if value == nil { + return "" + } + return fmt.Sprint(value) +} + +func optionalString(value string) *string { + value = strings.TrimSpace(value) + if value == "" { + return nil + } + return &value +} diff --git a/agents/drivers/argo-go/metadata_test.go b/agents/drivers/argo-go/metadata_test.go new file mode 100644 index 0000000000..0a83e1127e --- /dev/null +++ b/agents/drivers/argo-go/metadata_test.go @@ -0,0 +1,748 @@ +package main + +import ( + "context" + "database/sql/driver" + "encoding/json" + "errors" + "reflect" + "strings" + "testing" + + "github.com/t8y2/dbx/agents/go-common/gohive" +) + +func TestShowTablesRowName(t *testing.T) { + if value := showTablesRowName([]string{"database", "tableName", "isTemporary"}, []any{"default", "events", false}); value != "events" { + t.Fatalf("unexpected table name: %q", value) + } + if value := showTablesRowName([]string{"tab_name"}, []any{"fallback"}); value != "fallback" { + t.Fatalf("unexpected fallback table name: %q", value) + } +} + +// The argo agent always reports ArgoDB identity; Kyuubi/Impala/Hive remain on hive-go. +func TestConnectionInfoReportsArgoIdentity(t *testing.T) { + behavior := &scriptedBehavior{ + query: func(ctx context.Context, query string) (driver.Rows, error) { + switch query { + case "SELECT VERSION()": + return newScriptedRows(ctx, []string{"version"}, []string{"STRING"}, [][]driver.Value{{"3.5.8"}}), nil + case "SELECT CURRENT_USER()": + return newScriptedRows(ctx, []string{"current_user"}, []string{"STRING"}, [][]driver.Value{{"dbx"}}), nil + default: + return nil, errors.New("unexpected query: " + query) + } + }, + } + server := newScriptedServer(t, behavior) + server.params.DatabaseType = "argo" + server.config.Username = "fallback" + + info, err := server.connectionInfo() + if err != nil { + t.Fatal(err) + } + if info["compatibilityMode"] != "argo" || info["username"] != "dbx" || info["version"] != "3.5.8" { + t.Fatalf("unexpected Argo connection info: %#v", info) + } + databaseInfo, ok := info["databaseInfo"].(map[string]string) + if !ok || databaseInfo["productName"] != "ArgoDB (Transwarp)" || databaseInfo["driverName"] != "DBX ArgoDB Go Agent" { + t.Fatalf("unexpected Argo database identity: %#v", info["databaseInfo"]) + } +} + +func TestGetObjectSourceReturnsProtocolObject(t *testing.T) { + behavior := &scriptedBehavior{ + query: func(ctx context.Context, query string) (driver.Rows, error) { + if query != "SHOW CREATE TABLE `dbx_kyuubi_demo`.`high_value_orders`" { + t.Fatalf("unexpected query: %q", query) + } + return newScriptedRows( + ctx, + []string{"createtab_stmt"}, + []string{"STRING"}, + [][]driver.Value{ + {"CREATE VIEW dbx_kyuubi_demo.high_value_orders"}, + {"AS SELECT id, customer, amount FROM dbx_kyuubi_demo.orders WHERE amount >= 50"}, + }, + ), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + + result, _, err := server.dispatch("get_object_source", map[string]json.RawMessage{ + "schema": json.RawMessage(`"dbx_kyuubi_demo"`), + "name": json.RawMessage(`"high_value_orders"`), + "object_type": json.RawMessage(`"VIEW"`), + }) + if err != nil { + t.Fatal(err) + } + source, ok := result.(objectSource) + if !ok { + t.Fatalf("get_object_source returned %T instead of objectSource", result) + } + if source.Name != "high_value_orders" || source.ObjectType != "VIEW" || source.Schema == nil || *source.Schema != "dbx_kyuubi_demo" { + t.Fatalf("unexpected object source metadata: %#v", source) + } + expected := "CREATE VIEW dbx_kyuubi_demo.high_value_orders\nAS SELECT id, customer, amount FROM dbx_kyuubi_demo.orders WHERE amount >= 50\n" + if source.Source != expected { + t.Fatalf("unexpected object source DDL: %q", source.Source) + } +} + +func TestListDatabasesUsesShowDatabasesBeforeHiveServerMetadata(t *testing.T) { + behavior := &scriptedBehavior{ + query: func(ctx context.Context, sql string) (driver.Rows, error) { + if sql != "SHOW DATABASES" { + t.Fatalf("unexpected query: %q", sql) + } + return newScriptedRows(ctx, []string{"database_name"}, []string{"STRING"}, [][]driver.Value{{"warehouse"}, {"default"}, {"default"}}), nil + }, + getSchemas: func(_ context.Context, pattern string) (gohive.MetadataResult, error) { + t.Fatalf("metadata fallback must not run after SHOW DATABASES succeeds: %q", pattern) + return gohive.MetadataResult{}, nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + values, err := server.listDatabases() + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(values, []databaseInfo{{Name: "default"}, {Name: "warehouse"}}) { + t.Fatalf("unexpected databases: %#v", values) + } +} + +func TestListDatabasesFallsBackToHiveServerMetadata(t *testing.T) { + behavior := &scriptedBehavior{ + query: func(context.Context, string) (driver.Rows, error) { + return nil, errors.New("SHOW DATABASES unsupported") + }, + getSchemas: func(_ context.Context, pattern string) (gohive.MetadataResult, error) { + if pattern != "%" { + t.Fatalf("unexpected schema pattern: %q", pattern) + } + return metadataResult([]string{"TABLE_SCHEM", "TABLE_CATALOG"}, []driver.Value{"warehouse", ""}, []driver.Value{"default", ""}, []driver.Value{"default", ""}), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + values, err := server.listDatabases() + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(values, []databaseInfo{{Name: "default"}, {Name: "warehouse"}}) { + t.Fatalf("unexpected databases: %#v", values) + } +} + +func TestListSchemasHonorsVisibleSchemaFilter(t *testing.T) { + behavior := &scriptedBehavior{ + query: func(ctx context.Context, sql string) (driver.Rows, error) { + if sql != "SHOW DATABASES" { + t.Fatalf("unexpected query: %q", sql) + } + return newScriptedRows( + ctx, + []string{"database_name"}, + []string{"STRING"}, + [][]driver.Value{{"default"}, {"analytics"}, {"system"}}, + ), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + + values, err := server.listSchemas([]string{"analytics", "missing"}) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(values, []string{"analytics"}) { + t.Fatalf("unexpected visible schemas: %#v", values) + } + + values, err = server.listSchemas([]string{}) + if err != nil { + t.Fatal(err) + } + if len(values) != 0 { + t.Fatalf("explicit empty visible schema filter must hide all schemas: %#v", values) + } +} + +func TestListTablesPreservesViewTypeCommentAndWindow(t *testing.T) { + behavior := &scriptedBehavior{ + getTables: func(_ context.Context, schema, table string, tableTypes []string) (gohive.MetadataResult, error) { + if schema != "analytics" || table != "%" || !reflect.DeepEqual(tableTypes, []string{"TABLE", "VIEW", "MATERIALIZED VIEW"}) { + t.Fatalf("unexpected GetTables request: schema=%q table=%q types=%#v", schema, table, tableTypes) + } + return metadataResult( + []string{"TABLE_CAT", "TABLE_SCHEM", "TABLE_NAME", "TABLE_TYPE", "REMARKS"}, + []driver.Value{"", "analytics", "events", "TABLE", "event data"}, + []driver.Value{"", "analytics", "events_view", "VIEW", "view data"}, + []driver.Value{"", "analytics", "other", "TABLE", nil}, + ), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + values, err := server.listTables("analytics", metadataListConstraints{Filter: "events", Offset: 1, Limit: 1}) + if err != nil { + t.Fatal(err) + } + comment := "view data" + expected := []tableInfo{{Name: "events_view", TableType: "VIEW", Comment: &comment}} + if !reflect.DeepEqual(values, expected) { + t.Fatalf("unexpected tables: %#v", values) + } +} + +func TestGetColumnsUsesHiveServerMetadataFields(t *testing.T) { + behavior := &scriptedBehavior{ + getColumns: func(_ context.Context, schema, table, column string) (gohive.MetadataResult, error) { + if schema != "analytics" || table != "events" || column != "%" { + t.Fatalf("unexpected GetColumns request: %q %q %q", schema, table, column) + } + return metadataResult( + []string{"COLUMN_NAME", "TYPE_NAME", "COLUMN_SIZE", "DECIMAL_DIGITS", "NULLABLE", "REMARKS", "COLUMN_DEF"}, + []driver.Value{"name", "string", int64(255), nil, int64(1), "显示名称", "unknown"}, + []driver.Value{"amount", "decimal(18,2)", int64(18), int64(2), int64(0), nil, nil}, + ), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + values, err := server.getColumns("analytics", "events") + if err != nil { + t.Fatal(err) + } + if len(values) != 2 { + t.Fatalf("unexpected columns: %#v", values) + } + if values[0].Name != "name" || !values[0].IsNullable || values[0].CharacterMaximumLength == nil || *values[0].CharacterMaximumLength != 255 || values[0].ColumnDefault == nil || *values[0].ColumnDefault != "unknown" || values[0].Comment == nil || *values[0].Comment != "显示名称" { + t.Fatalf("unexpected string column: %#v", values[0]) + } + if values[1].Name != "amount" || values[1].IsNullable || values[1].NumericPrecision == nil || *values[1].NumericPrecision != 18 || values[1].NumericScale == nil || *values[1].NumericScale != 2 || values[1].CharacterMaximumLength != nil { + t.Fatalf("unexpected decimal column: %#v", values[1]) + } +} + +func TestGetColumnsPreservesChineseDescribeFallbackComments(t *testing.T) { + behavior := &scriptedBehavior{ + getColumns: func(context.Context, string, string, string) (gohive.MetadataResult, error) { + return gohive.MetadataResult{}, errors.New("metadata unavailable") + }, + query: func(ctx context.Context, query string) (driver.Rows, error) { + if query != "DESCRIBE `analytics`.`events`" { + t.Fatalf("unexpected DESCRIBE query: %q", query) + } + return newScriptedRows( + ctx, + []string{"col_name", "data_type", "comment"}, + []string{"STRING", "STRING", "STRING"}, + [][]driver.Value{{"name", "string", "显示名称"}, {"amount", "decimal(18,2)", "含税金额"}}, + ), nil + }, + } + server := newScriptedServer(t, behavior) + values, err := server.getColumns("analytics", "events") + if err != nil { + t.Fatal(err) + } + if len(values) != 2 || values[0].Comment == nil || *values[0].Comment != "显示名称" || values[1].Comment == nil || *values[1].Comment != "含税金额" { + t.Fatalf("DESCRIBE comments changed: %#v", values) + } +} + +func TestTableCommentAndTypeInfoUseHiveServerMetadata(t *testing.T) { + behavior := &scriptedBehavior{ + getTables: func(_ context.Context, schema, table string, tableTypes []string) (gohive.MetadataResult, error) { + return metadataResult( + []string{"TABLE_SCHEM", "TABLE_NAME", "TABLE_TYPE", "REMARKS"}, + []driver.Value{schema, table, "TABLE", "table comment"}, + ), nil + }, + getTypeInfo: func(context.Context) (gohive.MetadataResult, error) { + return metadataResult([]string{"TYPE_NAME"}, []driver.Value{"STRING"}, []driver.Value{"decimal"}, []driver.Value{"STRING"}), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + comment, err := server.getTableComment("analytics", "events") + if err != nil || comment == nil || *comment != "table comment" { + t.Fatalf("unexpected table comment: %v, %v", comment, err) + } + types, err := server.listDataTypes() + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(types, []string{"decimal", "string"}) { + t.Fatalf("unexpected data types: %#v", types) + } +} + +func TestListTablesFallsBackToShowTablesAndViews(t *testing.T) { + behavior := &scriptedBehavior{ + getTables: func(context.Context, string, string, []string) (gohive.MetadataResult, error) { + return gohive.MetadataResult{}, errors.New("metadata unsupported") + }, + query: func(ctx context.Context, query string) (driver.Rows, error) { + switch query { + case "SHOW TABLES IN `analytics`": + return newScriptedRows( + ctx, + []string{"tab_name"}, + []string{"STRING"}, + [][]driver.Value{{"events"}, {"shared_name"}}, + ), nil + case "SHOW VIEWS IN `analytics`": + return newScriptedRows( + ctx, + []string{"view_name"}, + []string{"STRING"}, + [][]driver.Value{{"events_view"}, {"shared_name"}}, + ), nil + default: + t.Fatalf("unexpected fallback query: %q", query) + return nil, errors.New("unexpected fallback query") + } + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + values, err := server.listTables("analytics", metadataListConstraints{}) + if err != nil { + t.Fatal(err) + } + expected := []tableInfo{ + {Name: "events", TableType: "TABLE"}, + {Name: "events_view", TableType: "VIEW"}, + {Name: "shared_name", TableType: "VIEW"}, + } + if !reflect.DeepEqual(values, expected) { + t.Fatalf("unexpected fallback tables: %#v", values) + } +} + +func TestListTablesKeepsShowTablesResultsWhenShowViewsIsUnsupported(t *testing.T) { + behavior := &scriptedBehavior{ + getTables: func(context.Context, string, string, []string) (gohive.MetadataResult, error) { + return gohive.MetadataResult{}, errors.New("metadata unsupported") + }, + query: func(ctx context.Context, query string) (driver.Rows, error) { + switch query { + case "SHOW TABLES IN `analytics`": + return newScriptedRows(ctx, []string{"tab_name"}, []string{"STRING"}, [][]driver.Value{{"events"}}), nil + case "SHOW VIEWS IN `analytics`": + return nil, errors.New("SHOW VIEWS is unsupported") + default: + t.Fatalf("unexpected fallback query: %q", query) + return nil, errors.New("unexpected fallback query") + } + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + + values, err := server.listTables("analytics", metadataListConstraints{}) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(values, []tableInfo{{Name: "events", TableType: "TABLE"}}) { + t.Fatalf("unexpected fallback tables: %#v", values) + } +} + +func TestListTablesReturnsNonCapabilityShowViewsError(t *testing.T) { + behavior := &scriptedBehavior{ + getTables: func(context.Context, string, string, []string) (gohive.MetadataResult, error) { + return gohive.MetadataResult{}, errors.New("metadata unsupported") + }, + query: func(ctx context.Context, query string) (driver.Rows, error) { + switch query { + case "SHOW TABLES IN `analytics`": + return newScriptedRows(ctx, []string{"tab_name"}, []string{"STRING"}, [][]driver.Value{{"events"}}), nil + case "SHOW VIEWS IN `analytics`": + return nil, errors.New("permission denied for SHOW VIEWS") + default: + t.Fatalf("unexpected fallback query: %q", query) + return nil, errors.New("unexpected fallback query") + } + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + + _, err := server.listTables("analytics", metadataListConstraints{}) + if err == nil || !strings.Contains(err.Error(), "SHOW VIEWS fallback failed: permission denied") { + t.Fatalf("unexpected mixed fallback error: %v", err) + } +} + +func TestShowViewsUnsupported(t *testing.T) { + tests := []struct { + name string + err error + unsupported bool + }{ + {name: "explicit unsupported", err: errors.New("SHOW VIEWS is unsupported"), unsupported: true}, + {name: "not supported", err: errors.New("SHOW VIEWS is not supported before Hive 2.2"), unsupported: true}, + {name: "old parser", err: errors.New("ParseException: syntax error at or near VIEWS"), unsupported: true}, + {name: "permission", err: errors.New("permission denied for SHOW VIEWS")}, + {name: "timeout", err: context.DeadlineExceeded}, + {name: "cancel", err: context.Canceled}, + {name: "authentication", err: errors.New("authentication failed")}, + {name: "unsupported authentication", err: errors.New("unsupported authentication mechanism")}, + {name: "transport", err: errors.New("transport is closed")}, + {name: "unsupported transport", err: errors.New("transport does not support SASL")}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + if actual := showViewsUnsupported(test.err); actual != test.unsupported { + t.Fatalf("showViewsUnsupported(%v) = %v, want %v", test.err, actual, test.unsupported) + } + }) + } +} + +func TestListTablesFallbackHonorsExplicitTableType(t *testing.T) { + behavior := &scriptedBehavior{ + getTables: func(context.Context, string, string, []string) (gohive.MetadataResult, error) { + return gohive.MetadataResult{}, errors.New("metadata unsupported") + }, + query: func(ctx context.Context, query string) (driver.Rows, error) { + if query != "SHOW TABLES IN `analytics`" { + t.Fatalf("unexpected fallback query: %q", query) + } + return newScriptedRows(ctx, []string{"tab_name"}, []string{"STRING"}, [][]driver.Value{{"events"}}), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + + values, err := server.listTables("analytics", metadataListConstraints{ObjectTypes: []string{"TABLE"}}) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(values, []tableInfo{{Name: "events", TableType: "TABLE"}}) { + t.Fatalf("unexpected fallback tables: %#v", values) + } +} + +func TestListViewsFallsBackToShowViews(t *testing.T) { + behavior := &scriptedBehavior{ + getTables: func(context.Context, string, string, []string) (gohive.MetadataResult, error) { + return gohive.MetadataResult{}, errors.New("metadata unsupported") + }, + query: func(ctx context.Context, query string) (driver.Rows, error) { + if query != "SHOW VIEWS IN `analytics`" { + t.Fatalf("unexpected fallback query: %q", query) + } + return newScriptedRows(ctx, []string{"view_name"}, []string{"STRING"}, [][]driver.Value{{"events_view"}}), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + values, err := server.listTables("analytics", metadataListConstraints{ObjectTypes: []string{"VIEW"}}) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(values, []tableInfo{{Name: "events_view", TableType: "VIEW"}}) { + t.Fatalf("unexpected fallback views: %#v", values) + } +} + +func TestListViewsReturnsFallbackErrorWhenShowViewsIsUnsupported(t *testing.T) { + behavior := &scriptedBehavior{ + getTables: func(context.Context, string, string, []string) (gohive.MetadataResult, error) { + return gohive.MetadataResult{}, errors.New("metadata unsupported") + }, + query: func(_ context.Context, query string) (driver.Rows, error) { + if query != "SHOW VIEWS IN `analytics`" { + t.Fatalf("unexpected fallback query: %q", query) + } + return nil, errors.New("SHOW VIEWS is unsupported") + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + + _, err := server.listTables("analytics", metadataListConstraints{ObjectTypes: []string{"VIEW"}}) + if err == nil || !strings.Contains(err.Error(), "SHOW VIEWS fallback failed") { + t.Fatalf("unexpected explicit view fallback error: %v", err) + } +} + +func metadataResult(columns []string, rows ...[]driver.Value) gohive.MetadataResult { + return gohive.MetadataResult{Columns: columns, Rows: rows} +} + +func TestListObjectsIncludesProceduresAndFunctionsFromSystemViews(t *testing.T) { + proceduresQuery := "SELECT procedure_name FROM system.procedures_v WHERE lower(database_name) = lower('ods') AND lower(procedure_name) LIKE '%sp%' ORDER BY procedure_name" + functionsQuery := "SELECT function_name FROM system.functions_v WHERE lower(database_name) = lower('ods') AND lower(function_name) LIKE '%sp%' ORDER BY function_name" + behavior := &scriptedBehavior{ + query: func(_ context.Context, query string) (driver.Rows, error) { + switch query { + case proceduresQuery: + return newScriptedRows(context.Background(), []string{"procedure_name"}, []string{"STRING"}, [][]driver.Value{ + {"sp_daily_etl"}, {"sp_hourly_agg"}, + }), nil + case functionsQuery: + return newScriptedRows(context.Background(), []string{"function_name"}, []string{"STRING"}, [][]driver.Value{ + {"fn_clean"}, + }), nil + default: + return nil, errors.New("unexpected query: " + query) + } + }, + } + server := newScriptedServer(t, behavior) + server.params.DatabaseType = "argo" + defer server.disconnect() + + values, err := server.listObjects("ods", "ods", metadataListConstraints{ + ObjectTypes: []string{"PROCEDURE", "FUNCTION"}, + Filter: "sp", // "sp" overlaps the procedure filter; both lists get "%sp%" applied + }) + if err != nil { + t.Fatalf("listObjects: %v", err) + } + t.Logf("values: %+v", values) + // listObjects applies the same filter pattern to procedures and functions; + // we passed "sp" so procedures match and functions don't (because the function + // list returns rows regardless of filter — that is, listObjects calls each + // listRoutines call with the same Filter). Verify both queries ran and + // the procedure name is present. + queries, _, _, _ := behavior.snapshot() + t.Logf("queries count: %d", len(queries)) + for i, q := range queries { + t.Logf("query[%d]: %q", i, q) + } + if len(queries) != 2 { + t.Fatalf("expected exactly 2 routine queries, got %d: %v", len(queries), queries) + } + if values == nil { + t.Fatal("expected non-nil object list") + } + names := make([]string, len(values)) + for i, v := range values { + names[i] = v.Name + ":" + v.ObjectType + } + joined := strings.Join(names, ",") + if !strings.Contains(joined, "sp_daily_etl:PROCEDURE") { + t.Fatalf("expected procedures in result, got %v", values) + } + if !strings.Contains(joined, "fn_clean:FUNCTION") { + t.Fatalf("expected functions in result, got %v", values) + } +} + +func TestGetObjectSourceRoutesProceduresToSystemProceduresView(t *testing.T) { + expectedSQL := "SELECT full_text FROM system.procedures_v WHERE lower(database_name) = lower('ods') AND procedure_name = 'sp_daily_etl'" + behavior := &scriptedBehavior{ + query: func(_ context.Context, query string) (driver.Rows, error) { + if query != expectedSQL { + return nil, errors.New("unexpected query: " + query) + } + return newScriptedRows(context.Background(), []string{"full_text"}, []string{"STRING"}, [][]driver.Value{ + {"-- daily ETL pipeline"}, + {"INSERT OVERWRITE TABLE ods.daily_summary SELECT * FROM staging.events"}, + }), nil + }, + } + server := newScriptedServer(t, behavior) + server.params.DatabaseType = "argo" + defer server.disconnect() + + result, _, err := server.dispatch("get_object_source", map[string]json.RawMessage{ + "schema": json.RawMessage(`"ods"`), + "name": json.RawMessage(`"sp_daily_etl"`), + "object_type": json.RawMessage(`"PROCEDURE"`), + }) + if err != nil { + t.Fatal(err) + } + source, ok := result.(objectSource) + if !ok { + t.Fatalf("get_object_source returned %T instead of objectSource", result) + } + if source.Name != "sp_daily_etl" || source.ObjectType != "PROCEDURE" { + t.Fatalf("unexpected object source metadata: %#v", source) + } + expected := "-- daily ETL pipeline\nINSERT OVERWRITE TABLE ods.daily_summary SELECT * FROM staging.events\n" + if source.Source != expected { + t.Fatalf("unexpected procedure source: %q", source.Source) + } +} + +func TestGetObjectSourceRoutesFunctionsToSystemFunctionsView(t *testing.T) { + expectedSQL := "SELECT full_text FROM system.functions_v WHERE lower(database_name) = lower('ods') AND function_name = 'fn_clean'" + behavior := &scriptedBehavior{ + query: func(_ context.Context, query string) (driver.Rows, error) { + if query != expectedSQL { + return nil, errors.New("unexpected query: " + query) + } + return newScriptedRows(context.Background(), []string{"full_text"}, []string{"STRING"}, [][]driver.Value{ + {"-- cleanup helper"}, + }), nil + }, + } + server := newScriptedServer(t, behavior) + server.params.DatabaseType = "argo" + defer server.disconnect() + + result, _, err := server.dispatch("get_object_source", map[string]json.RawMessage{ + "schema": json.RawMessage(`"ods"`), + "name": json.RawMessage(`"fn_clean"`), + "object_type": json.RawMessage(`"FUNCTION"`), + }) + if err != nil { + t.Fatal(err) + } + source, ok := result.(objectSource) + if !ok { + t.Fatalf("get_object_source returned %T instead of objectSource", result) + } + if source.Source != "-- cleanup helper\n" { + t.Fatalf("unexpected function source: %q", source.Source) + } +} + +func TestListRoutinesUsesDatabaseParameterWhenSchemaEmpty(t *testing.T) { + proceduresQuery := "SELECT procedure_name FROM system.procedures_v WHERE lower(database_name) = lower('ods') AND lower(procedure_name) LIKE '%%' ORDER BY procedure_name" + defaultQuery := "SELECT procedure_name FROM system.procedures_v WHERE lower(database_name) = lower('default') AND lower(procedure_name) LIKE '%%' ORDER BY procedure_name" + behavior := &scriptedBehavior{ + query: func(_ context.Context, query string) (driver.Rows, error) { + switch query { + case defaultQuery: + return newScriptedRows(context.Background(), []string{"procedure_name"}, []string{"STRING"}, [][]driver.Value{}), nil + case proceduresQuery: + return newScriptedRows(context.Background(), []string{"procedure_name"}, []string{"STRING"}, [][]driver.Value{ + {"sp_daily_etl"}, + }), nil + default: + return nil, errors.New("unexpected query: " + query) + } + }, + } + server := newScriptedServer(t, behavior) + server.params.DatabaseType = "argo" + server.config.Database = "default" + defer server.disconnect() + + values, err := server.listObjects("ods", "", metadataListConstraints{ + ObjectTypes: []string{"PROCEDURE"}, + }) + if err != nil { + t.Fatalf("listObjects: %v", err) + } + if len(values) != 1 || values[0].Name != "sp_daily_etl" { + t.Fatalf("expected procedure from database parameter, got %+v", values) + } +} + +func TestListRoutinesDoesNotFallbackToConnectionDefaultForExplicitSchema(t *testing.T) { + defaultQuery := "SELECT procedure_name FROM system.procedures_v WHERE lower(database_name) = lower('default') AND lower(procedure_name) LIKE '%%' ORDER BY procedure_name" + behavior := &scriptedBehavior{ + query: func(_ context.Context, query string) (driver.Rows, error) { + if query == defaultQuery { + // Serving this row proves the driver fell back to the + // connection default after the explicit schema came back empty. + return newScriptedRows(context.Background(), []string{"procedure_name"}, []string{"STRING"}, [][]driver.Value{ + {"sp_should_not_leak"}, + }), nil + } + return newScriptedRows(context.Background(), []string{"procedure_name"}, []string{"STRING"}, [][]driver.Value{}), nil + }, + } + server := newScriptedServer(t, behavior) + server.params.DatabaseType = "argo" + server.config.Database = "default" + defer server.disconnect() + + values, err := server.listObjects("", "ods", metadataListConstraints{ + ObjectTypes: []string{"PROCEDURE"}, + }) + if err != nil { + t.Fatalf("listObjects: %v", err) + } + if len(values) != 0 { + t.Fatalf("expected no routines to leak from the connection default database, got %+v", values) + } +} + +func TestGetObjectSourceUsesDatabaseParameterForRoutineSource(t *testing.T) { + expectedSQL := "SELECT full_text FROM system.procedures_v WHERE lower(database_name) = lower('ods') AND procedure_name = 'sp_daily_etl'" + behavior := &scriptedBehavior{ + query: func(_ context.Context, query string) (driver.Rows, error) { + if query != expectedSQL { + return nil, errors.New("unexpected query: " + query) + } + return newScriptedRows(context.Background(), []string{"full_text"}, []string{"STRING"}, [][]driver.Value{ + {"CREATE PROCEDURE sp_daily_etl() BEGIN SELECT 1; END"}, + }), nil + }, + } + server := newScriptedServer(t, behavior) + server.params.DatabaseType = "argo" + server.config.Database = "default" + defer server.disconnect() + + result, _, err := server.dispatch("get_object_source", map[string]json.RawMessage{ + "database": json.RawMessage(`"ods"`), + "schema": json.RawMessage(`""`), + "name": json.RawMessage(`"sp_daily_etl"`), + "object_type": json.RawMessage(`"PROCEDURE"`), + }) + if err != nil { + t.Fatal(err) + } + source, ok := result.(objectSource) + if !ok { + t.Fatalf("get_object_source returned %T instead of objectSource", result) + } + if source.Source != "CREATE PROCEDURE sp_daily_etl() BEGIN SELECT 1; END\n" { + t.Fatalf("unexpected procedure source: %q", source.Source) + } +} + +// Unlike hive-go (which gates routine listing on the connection's database_type so +// vanilla Apache Hive never fires the catalog-view queries), the argo agent serves +// 星环Argo (Transwarp ArgoDB) exclusively and must always query the routine views. +func TestListObjectsQueriesRoutineViewsUnconditionally(t *testing.T) { + behavior := &scriptedBehavior{ + query: func(ctx context.Context, query string) (driver.Rows, error) { + switch { + case strings.Contains(strings.ToUpper(query), "PROCEDURES_V"): + return newScriptedRows(ctx, []string{"procedure_name"}, []string{"STRING"}, [][]driver.Value{{"sp_etl_log"}}), nil + case strings.Contains(strings.ToUpper(query), "FUNCTIONS_V"): + return newScriptedRows(ctx, []string{"function_name"}, []string{"STRING"}, nil), nil + default: + return newScriptedRows(ctx, []string{"tab_name"}, []string{"STRING"}, [][]driver.Value{{"events"}}), nil + } + }, + } + server := newScriptedServer(t, behavior) + server.params.DatabaseType = "argo" + defer server.disconnect() + + values, err := server.listObjects("ods", "ods", metadataListConstraints{ + ObjectTypes: []string{"PROCEDURE", "FUNCTION"}, + }) + if err != nil { + t.Fatalf("listObjects: %v", err) + } + found := false + for _, v := range values { + if v.ObjectType == "PROCEDURE" && v.Name == "sp_etl_log" { + found = true + } + } + if !found { + t.Fatalf("argo agent must surface routines regardless of database_type, got: %+v", values) + } +} diff --git a/agents/drivers/argo-go/protocol_error.go b/agents/drivers/argo-go/protocol_error.go new file mode 100644 index 0000000000..76cfb36e1a --- /dev/null +++ b/agents/drivers/argo-go/protocol_error.go @@ -0,0 +1,145 @@ +package main + +import ( + "context" + "database/sql" + "errors" + "fmt" + "io" + "net" + "strings" + + "github.com/t8y2/dbx/agents/go-common/gohive" +) + +type rpcError struct { + Code int `json:"code"` + Message string `json:"message"` + Data *rpcErrorData `json:"data,omitempty"` +} + +type rpcErrorData struct { + Category string `json:"category"` + Retryable bool `json:"retryable"` + SessionDisposition string `json:"sessionDisposition"` + Stage string `json:"stage"` + ContractVersion int `json:"contractVersion"` + OperationOutcome string `json:"operationOutcome"` + SQLState string `json:"sqlState,omitempty"` + VendorCode int32 `json:"vendorCode,omitempty"` + ExceptionClass string `json:"exceptionClass,omitempty"` + AgentSessionID string `json:"agentSessionId,omitempty"` +} + +func classifyRPCError(method, agentSessionID string, err error) *rpcError { + stage := rpcErrorStage(method) + data := &rpcErrorData{ + Category: "protocol", + Retryable: false, + SessionDisposition: "keep", + Stage: stage, + ContractVersion: 1, + OperationOutcome: rpcOperationOutcome(stage), + ExceptionClass: safeRPCDiagnostic(fmt.Sprintf("%T", err), 160), + AgentSessionID: strings.TrimSpace(agentSessionID), + } + + var hiveError *gohive.Error + if errors.As(err, &hiveError) { + data.Category = "sql" + data.SQLState = safeRPCDiagnostic(hiveError.SQLState, 160) + data.VendorCode = int32(hiveError.ErrorCode) + data.ExceptionClass = "gohive.Error" + } else if errors.Is(err, context.Canceled) { + data.Category = "canceled" + data.SessionDisposition = "quarantine" + } else if errors.Is(err, context.DeadlineExceeded) || isTimeoutError(err) { + data.Category = "timeout" + data.Retryable = stage == "connect" || stage == "validate" + data.SessionDisposition = "quarantine" + } else if isConnectionError(err) { + data.Category = "connection" + data.Retryable = stage == "connect" || stage == "validate" + if stage != "connect" { + data.SessionDisposition = "quarantine" + } + } else if isHiveSQLError(err) { + data.Category = "sql" + } + return &rpcError{Code: -1, Message: err.Error(), Data: data} +} + +func rpcErrorStage(method string) string { + switch method { + case "connect", "open_session", "test_connection": + return "connect" + case "validate_connection", "validate_session": + return "validate" + case "cancel_session": + return "cancel" + case "close_session", "disconnect", "close_query_session", "close_table_read_session", "shutdown": + return "close" + case "fetch_query_page", "fetch_table_read_page": + return "fetch" + case "handshake", "": + return "request" + default: + return "execute" + } +} + +func rpcOperationOutcome(stage string) string { + if stage == "request" || stage == "connect" || stage == "validate" { + return "not_started" + } + return "unknown" +} + +func isTimeoutError(err error) bool { + var timeout interface{ Timeout() bool } + return errors.As(err, &timeout) && timeout.Timeout() +} + +func isConnectionError(err error) bool { + if errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) || errors.Is(err, sql.ErrConnDone) { + return true + } + var networkError *net.OpError + if errors.As(err, &networkError) { + return true + } + lower := strings.ToLower(err.Error()) + for _, marker := range []string{ + "connection refused", "connection reset", "broken pipe", "connection closed", "connection lost", + "unexpected eof", "no route to host", "all hiveserver2 endpoints failed", "transport is not open", + "socket is closed", "not open", + } { + if strings.Contains(lower, marker) { + return true + } + } + return false +} + +func isHiveSQLError(err error) bool { + lower := strings.ToLower(err.Error()) + for _, marker := range []string{"semanticexception", "parseexception", "hive error", "sqlstate", "error while compiling statement"} { + if strings.Contains(lower, marker) { + return true + } + } + return false +} + +func safeRPCDiagnostic(value string, maxLength int) string { + var result strings.Builder + for _, char := range value { + if result.Len() >= maxLength { + break + } + if char >= 0x21 && char <= 0x7e { + result.WriteRune(char) + } + } + return result.String() +} diff --git a/agents/drivers/argo-go/query.go b/agents/drivers/argo-go/query.go new file mode 100644 index 0000000000..734d4fa967 --- /dev/null +++ b/agents/drivers/argo-go/query.go @@ -0,0 +1,529 @@ +package main + +import ( + "context" + "database/sql" + "database/sql/driver" + "encoding/json" + "errors" + "fmt" + "strings" + "time" + + "github.com/t8y2/dbx/agents/go-common/gohive" +) + +func (server *server) validateConnection() error { + connection, err := server.requireConnection() + if err != nil { + return err + } + ctx, cancel := context.WithTimeout(context.Background(), server.config.ConnectTimeout) + defer cancel() + return connection.PingContext(ctx) +} + +func (server *server) executeQuery(options queryOptions) (queryResult, error) { + started := time.Now() + if options.FetchSize <= 0 { + options.FetchSize = server.effectiveFetchSize() + } + sqlText := trimStatementSQL(options.SQL) + if sqlText == "" { + return queryResult{}, errors.New("SQL is required") + } + maxRows := options.MaxRows + if maxRows <= 0 { + maxRows = defaultMaxRows + } + connection, err := server.requireConnection() + if err != nil { + return queryResult{}, err + } + + ctx, cancel := queryContext(options.TimeoutSecs) + server.setActiveOperation(cancel) + defer server.clearActiveOperation(cancel) + if err := server.applySchemaContext(ctx, connection, effectiveSchema(options)); err != nil { + return queryResult{}, err + } + rows, affected, hasResultSet, err := executeHiveStatement(ctx, connection, sqlText, options.FetchSize) + if err != nil { + return queryResult{}, err + } + if !hasResultSet { + return queryResult{ + Columns: []string{}, + ColumnTypes: []string{}, + Rows: [][]any{}, + AffectedRows: affected, + ExecutionTimeMS: time.Since(started).Milliseconds(), + Truncated: false, + }, nil + } + defer rows.Close() + columns, columnTypes, err := queryColumnMetadata(rows) + if err != nil { + return queryResult{}, err + } + values, truncated, err := readSQLRows(rows, columnTypes, maxRows) + if err != nil { + return queryResult{}, err + } + return queryResult{ + Columns: columns, + ColumnTypes: columnTypes, + Rows: values, + AffectedRows: 0, + ExecutionTimeMS: time.Since(started).Milliseconds(), + Truncated: truncated, + }, nil +} + +func (server *server) executeQueryPage(options queryOptions, requestedPageSize int) (queryPageResult, error) { + started := time.Now() + if options.FetchSize <= 0 { + options.FetchSize = server.effectiveFetchSize() + } + server.expireIdleQuerySessions(started) + sqlText := trimStatementSQL(options.SQL) + if sqlText == "" { + return queryPageResult{}, errors.New("SQL is required") + } + pageSize := requestedPageSize + if pageSize <= 0 { + pageSize = defaultPageSize + } + maxRows := options.MaxRows + if maxRows <= 0 { + maxRows = defaultMaxRows + } + connection, err := server.requireConnection() + if err != nil { + return queryPageResult{}, err + } + ctx, cancel := queryContext(options.TimeoutSecs) + server.setActiveOperation(cancel) + if err := server.applySchemaContext(ctx, connection, effectiveSchema(options)); err != nil { + server.clearActiveOperation(cancel) + return queryPageResult{}, err + } + rows, affected, hasResultSet, err := executeHiveStatement(ctx, connection, sqlText, options.FetchSize) + if err != nil { + server.clearActiveOperation(cancel) + return queryPageResult{}, err + } + if !hasResultSet { + server.clearActiveOperation(cancel) + return queryPageResult{ + Columns: []string{}, + ColumnTypes: []string{}, + Rows: [][]any{}, + AffectedRows: affected, + ExecutionTimeMS: time.Since(started).Milliseconds(), + Truncated: false, + SessionID: nil, + HasMore: false, + }, nil + } + + columns, columnTypes, err := queryColumnMetadata(rows) + if err != nil { + _ = rows.Close() + server.clearActiveOperation(cancel) + return queryPageResult{}, err + } + server.nextSessionID++ + sessionID := fmt.Sprintf("hive-%d", server.nextSessionID) + state := &querySession{ + rows: rows, + columns: columns, + columnTypes: columnTypes, + remaining: maxRows, + cancel: cancel, + lastAccessed: started, + } + server.querySessions[sessionID] = state + page, hasMore, truncated, err := server.readQuerySessionPage(ctx, state, pageSize) + server.activeMu.Lock() + if server.activeCancel != nil { + server.activeCancel = nil + } + server.activeMu.Unlock() + if err != nil { + server.closeQuerySession(sessionID) + return queryPageResult{}, err + } + if !hasMore { + server.closeQuerySession(sessionID) + return queryPageResult{ + Columns: columns, + ColumnTypes: columnTypes, + Rows: page, + ExecutionTimeMS: time.Since(started).Milliseconds(), + Truncated: truncated, + SessionID: nil, + HasMore: false, + }, nil + } + return queryPageResult{ + Columns: columns, + ColumnTypes: columnTypes, + Rows: page, + ExecutionTimeMS: time.Since(started).Milliseconds(), + Truncated: false, + SessionID: &sessionID, + HasMore: true, + }, nil +} + +func (server *server) fetchQueryPage(sessionID string, requestedPageSize int) (queryPageResult, error) { + server.expireIdleQuerySessions(time.Now()) + state := server.querySessions[sessionID] + if state == nil { + return queryPageResult{}, errors.New("query session not found") + } + pageSize := requestedPageSize + if pageSize <= 0 { + pageSize = defaultPageSize + } + ctx := context.Background() + server.setActiveOperation(state.cancel) + page, hasMore, truncated, err := server.readQuerySessionPage(ctx, state, pageSize) + server.activeMu.Lock() + server.activeCancel = nil + server.activeMu.Unlock() + if err != nil { + server.closeQuerySession(sessionID) + return queryPageResult{}, err + } + var resultSessionID *string + if hasMore { + resultSessionID = &sessionID + } else { + server.closeQuerySession(sessionID) + } + return queryPageResult{ + Columns: state.columns, + ColumnTypes: state.columnTypes, + Rows: page, + Truncated: truncated, + SessionID: resultSessionID, + HasMore: hasMore, + }, nil +} + +func (server *server) readQuerySessionPage(ctx context.Context, state *querySession, pageSize int) ([][]any, bool, bool, error) { + state.lastAccessed = time.Now() + values := make([][]any, 0, min(pageSize, state.remaining)) + if state.pending != nil && state.remaining > 0 { + values = append(values, state.pending) + state.pending = nil + state.remaining-- + } + for len(values) < pageSize && state.remaining > 0 { + row, ok, err := nextSQLRow(state.rows, state.columnTypes) + if err != nil { + return nil, false, false, err + } + if !ok { + return values, false, false, nil + } + values = append(values, row) + state.remaining-- + select { + case <-ctx.Done(): + return nil, false, false, ctx.Err() + default: + } + } + if state.remaining == 0 { + _, ok, err := nextSQLRow(state.rows, state.columnTypes) + if err != nil { + return nil, false, false, err + } + return values, false, ok, nil + } + row, ok, err := nextSQLRow(state.rows, state.columnTypes) + if err != nil { + return nil, false, false, err + } + if !ok { + return values, false, false, nil + } + state.pending = row + return values, true, false, nil +} + +func (server *server) closeQuerySession(sessionID string) bool { + state := server.querySessions[sessionID] + if state == nil { + return false + } + delete(server.querySessions, sessionID) + state.cancel() + _ = state.rows.Close() + return true +} + +func (server *server) closeAllQuerySessions() error { + var failures []string + for sessionID, state := range server.querySessions { + delete(server.querySessions, sessionID) + state.cancel() + if err := state.rows.Close(); err != nil { + failures = append(failures, fmt.Sprintf("%s: %v", sessionID, err)) + } + } + if len(failures) > 0 { + return errors.New(strings.Join(failures, "; ")) + } + return nil +} + +func (server *server) expireIdleQuerySessions(now time.Time) int { + expired := make([]string, 0) + for sessionID, state := range server.querySessions { + if !state.lastAccessed.IsZero() && now.Sub(state.lastAccessed) >= querySessionIdleTime { + expired = append(expired, sessionID) + } + } + for _, sessionID := range expired { + server.closeQuerySession(sessionID) + } + return len(expired) +} + +func (server *server) executeStatements(params map[string]json.RawMessage, transaction bool) (queryResult, error) { + started := time.Now() + statements := stringSliceParam(params, "statements") + if len(statements) == 0 { + return queryResult{}, errors.New("statements are required") + } + connection, err := server.requireConnection() + if err != nil { + return queryResult{}, err + } + ctx, cancel := queryContext(intParam(params, "timeoutSecs")) + server.setActiveOperation(cancel) + defer server.clearActiveOperation(cancel) + if err := server.applySchemaContext(ctx, connection, firstNonEmpty(stringParam(params, "schema"), stringParam(params, "database"))); err != nil { + return queryResult{}, err + } + + var affected int64 + if transaction { + tx, beginErr := connection.BeginTx(ctx, nil) + if beginErr == nil { + for _, statement := range statements { + trimmed := trimStatementSQL(statement) + if trimmed == "" { + continue + } + result, execErr := tx.ExecContext(ctx, trimmed) + if execErr != nil { + _ = tx.Rollback() + return queryResult{}, execErr + } + count, _ := result.RowsAffected() + affected += max(count, 0) + } + if err := tx.Commit(); err != nil { + return queryResult{}, err + } + return emptyQueryResult(affected, started), nil + } + if !transactionUnsupported(beginErr) { + return queryResult{}, beginErr + } + } + for _, statement := range statements { + trimmed := trimStatementSQL(statement) + if trimmed == "" { + continue + } + result, err := connection.ExecContext(ctx, trimmed) + if err != nil { + return queryResult{}, err + } + count, _ := result.RowsAffected() + affected += max(count, 0) + } + return emptyQueryResult(affected, started), nil +} + +func transactionUnsupported(err error) bool { + if err == nil { + return false + } + if errors.Is(err, sql.ErrTxDone) { + return false + } + message := strings.ToLower(err.Error()) + return errors.Is(err, driver.ErrSkip) || + strings.Contains(message, "transactions are not supported") || + strings.Contains(message, "transaction is not supported") || + strings.Contains(message, "unsupported transaction") || + strings.Contains(message, "driver: skip fast-path") +} + +func emptyQueryResult(affected int64, started time.Time) queryResult { + return queryResult{ + Columns: []string{}, + ColumnTypes: []string{}, + Rows: [][]any{}, + AffectedRows: affected, + ExecutionTimeMS: time.Since(started).Milliseconds(), + Truncated: false, + } +} + +func (server *server) applySchemaContext(ctx context.Context, connection *sql.Conn, schema string) error { + schema = strings.TrimSpace(schema) + if schema == "" || strings.EqualFold(schema, server.config.Database) { + return nil + } + _, err := connection.ExecContext(ctx, "USE "+quoteHiveIdentifier(schema)) + return err +} + +func queryColumnMetadata(rows *sql.Rows) ([]string, []string, error) { + columns, err := rows.Columns() + if err != nil { + return nil, nil, err + } + types, err := rows.ColumnTypes() + if err != nil { + return nil, nil, err + } + columnTypes := make([]string, len(columns)) + for index := range columns { + if index < len(types) { + columnTypes[index] = strings.ToLower(strings.TrimSpace(types[index].DatabaseTypeName())) + } + } + return columns, columnTypes, nil +} + +func readSQLRows(rows *sql.Rows, columnTypes []string, limit int) ([][]any, bool, error) { + values := make([][]any, 0, min(limit, defaultFetchSize)) + for len(values) < limit { + row, ok, err := nextSQLRow(rows, columnTypes) + if err != nil { + return nil, false, err + } + if !ok { + return values, false, nil + } + values = append(values, row) + } + _, ok, err := nextSQLRow(rows, columnTypes) + return values, ok, err +} + +func nextSQLRow(rows *sql.Rows, columnTypes []string) ([]any, bool, error) { + if !rows.Next() { + if err := rows.Err(); err != nil { + return nil, false, err + } + return nil, false, nil + } + values := make([]any, len(columnTypes)) + targets := make([]any, len(values)) + for index := range values { + targets[index] = &values[index] + } + if err := rows.Scan(targets...); err != nil { + return nil, false, err + } + for index, value := range values { + values[index] = normalizeHiveValue(value, columnTypes[index]) + } + return values, true, nil +} + +func normalizeHiveValue(value any, columnType string) any { + if value == nil { + return nil + } + switch typed := value.(type) { + case []byte: + return bytesToHex(typed) + case time.Time: + return formatHiveJDBCDateTime(typed, columnType) + case fmt.Stringer: + return typed.String() + case string: + return typed + default: + return fmt.Sprint(value) + } +} + +func formatHiveJDBCDateTime(value time.Time, columnType string) string { + if strings.EqualFold(strings.TrimSpace(columnType), "DATE") { + return value.Format("2006-01-02") + } + base := value.Format("2006-01-02 15:04:05") + if value.Nanosecond() == 0 { + return base + ".0" + } + fraction := strings.TrimRight(fmt.Sprintf("%09d", value.Nanosecond()), "0") + return base + "." + fraction +} + +func executeHiveStatement( + ctx context.Context, + connection *sql.Conn, + sqlText string, + fetchSize int, +) (*sql.Rows, int64, bool, error) { + ctx = gohive.WithFetchSize(ctx, fetchSize) + rows, err := connection.QueryContext(ctx, sqlText) + if err == nil { + return rows, 0, true, nil + } + var nonQuery *gohive.NonQueryResult + if errors.As(err, &nonQuery) { + return nil, max(nonQuery.AffectedRows, 0), false, nil + } + return nil, 0, false, err +} + +func bytesToHex(value []byte) string { + const digits = "0123456789abcdef" + result := make([]byte, 2+len(value)*2) + result[0] = '0' + result[1] = 'x' + for index, current := range value { + result[2+index*2] = digits[current>>4] + result[3+index*2] = digits[current&0x0f] + } + return string(result) +} + +func queryContext(timeoutSecs int) (context.Context, context.CancelFunc) { + if timeoutSecs > 0 { + return context.WithTimeout(context.Background(), time.Duration(timeoutSecs)*time.Second) + } + return context.WithCancel(context.Background()) +} + +func (server *server) effectiveFetchSize() int { + if server.config.FetchSize > 0 { + return server.config.FetchSize + } + return defaultFetchSize +} + +func effectiveSchema(options queryOptions) string { + return firstNonEmpty(options.Schema, options.Database) +} + +func trimStatementSQL(sqlText string) string { + return strings.TrimSpace(strings.TrimSuffix(strings.TrimSpace(sqlText), ";")) +} + +func quoteHiveIdentifier(value string) string { + return "`" + strings.ReplaceAll(value, "`", "``") + "`" +} diff --git a/agents/drivers/argo-go/query_test.go b/agents/drivers/argo-go/query_test.go new file mode 100644 index 0000000000..e49cdd5e46 --- /dev/null +++ b/agents/drivers/argo-go/query_test.go @@ -0,0 +1,623 @@ +package main + +import ( + "context" + "database/sql" + "database/sql/driver" + "encoding/json" + "errors" + "fmt" + "io" + "reflect" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/t8y2/dbx/agents/go-common/gohive" +) + +var scriptedDriverSequence atomic.Uint64 + +type scriptedBehavior struct { + mu sync.Mutex + query func(context.Context, string) (driver.Rows, error) + exec func(context.Context, string) (driver.Result, error) + getSchemas func(context.Context, string) (gohive.MetadataResult, error) + getTables func(context.Context, string, string, []string) (gohive.MetadataResult, error) + getColumns func(context.Context, string, string, string) (gohive.MetadataResult, error) + getTypeInfo func(context.Context) (gohive.MetadataResult, error) + beginErr error + queries []string + executions []string + beginCalls int + closeCalls int +} + +func (behavior *scriptedBehavior) queryContext(ctx context.Context, query string) (driver.Rows, error) { + behavior.mu.Lock() + behavior.queries = append(behavior.queries, query) + operation := behavior.query + behavior.mu.Unlock() + if operation == nil { + return nil, fmt.Errorf("unexpected query: %s", query) + } + return operation(ctx, query) +} + +func (behavior *scriptedBehavior) execContext(ctx context.Context, query string) (driver.Result, error) { + behavior.mu.Lock() + behavior.executions = append(behavior.executions, query) + operation := behavior.exec + behavior.mu.Unlock() + if operation == nil { + return nil, fmt.Errorf("unexpected execution: %s", query) + } + return operation(ctx, query) +} + +func (behavior *scriptedBehavior) snapshot() (queries, executions []string, beginCalls, closeCalls int) { + behavior.mu.Lock() + defer behavior.mu.Unlock() + return append([]string(nil), behavior.queries...), append([]string(nil), behavior.executions...), behavior.beginCalls, behavior.closeCalls +} + +type scriptedDriver struct { + behavior *scriptedBehavior +} + +func (driverValue *scriptedDriver) Open(string) (driver.Conn, error) { + return &scriptedConnection{behavior: driverValue.behavior}, nil +} + +type scriptedConnection struct { + behavior *scriptedBehavior +} + +func (connection *scriptedConnection) Prepare(string) (driver.Stmt, error) { + return nil, errors.New("prepared statements are not supported") +} + +func (connection *scriptedConnection) Close() error { + connection.behavior.mu.Lock() + connection.behavior.closeCalls++ + connection.behavior.mu.Unlock() + return nil +} + +func (connection *scriptedConnection) Begin() (driver.Tx, error) { + connection.behavior.mu.Lock() + connection.behavior.beginCalls++ + connection.behavior.mu.Unlock() + if connection.behavior.beginErr != nil { + return nil, connection.behavior.beginErr + } + return nil, driver.ErrSkip +} + +func (connection *scriptedConnection) BeginTx(context.Context, driver.TxOptions) (driver.Tx, error) { + connection.behavior.mu.Lock() + connection.behavior.beginCalls++ + connection.behavior.mu.Unlock() + if connection.behavior.beginErr != nil { + return nil, connection.behavior.beginErr + } + return nil, driver.ErrSkip +} + +func (connection *scriptedConnection) Ping(context.Context) error { + return nil +} + +func (connection *scriptedConnection) QueryContext( + ctx context.Context, + query string, + _ []driver.NamedValue, +) (driver.Rows, error) { + return connection.behavior.queryContext(ctx, query) +} + +func (connection *scriptedConnection) ExecContext( + ctx context.Context, + query string, + _ []driver.NamedValue, +) (driver.Result, error) { + return connection.behavior.execContext(ctx, query) +} + +func (connection *scriptedConnection) GetHiveSchemas(ctx context.Context, pattern string) (gohive.MetadataResult, error) { + if connection.behavior.getSchemas == nil { + return gohive.MetadataResult{}, errors.New("GetSchemas unavailable") + } + return connection.behavior.getSchemas(ctx, pattern) +} + +func (connection *scriptedConnection) GetHiveTables(ctx context.Context, schema, table string, tableTypes []string) (gohive.MetadataResult, error) { + if connection.behavior.getTables == nil { + return gohive.MetadataResult{}, errors.New("GetTables unavailable") + } + return connection.behavior.getTables(ctx, schema, table, tableTypes) +} + +func (connection *scriptedConnection) GetHiveColumns(ctx context.Context, schema, table, column string) (gohive.MetadataResult, error) { + if connection.behavior.getColumns == nil { + return gohive.MetadataResult{}, errors.New("GetColumns unavailable") + } + return connection.behavior.getColumns(ctx, schema, table, column) +} + +func (connection *scriptedConnection) GetHiveTypeInfo(ctx context.Context) (gohive.MetadataResult, error) { + if connection.behavior.getTypeInfo == nil { + return gohive.MetadataResult{}, errors.New("GetTypeInfo unavailable") + } + return connection.behavior.getTypeInfo(ctx) +} + +type scriptedRows struct { + ctx context.Context + columns []string + types []string + values [][]driver.Value + blockAfter int + blocked chan struct{} + blockOnce sync.Once + + mu sync.Mutex + index int + closed bool +} + +func newScriptedRows(ctx context.Context, columns, types []string, values [][]driver.Value) *scriptedRows { + return &scriptedRows{ + ctx: ctx, + columns: columns, + types: types, + values: values, + blockAfter: -1, + } +} + +func (rows *scriptedRows) Columns() []string { + return append([]string(nil), rows.columns...) +} + +func (rows *scriptedRows) Close() error { + rows.mu.Lock() + rows.closed = true + rows.mu.Unlock() + return nil +} + +func (rows *scriptedRows) Next(destination []driver.Value) error { + rows.mu.Lock() + if rows.closed { + rows.mu.Unlock() + return io.EOF + } + index := rows.index + if rows.blockAfter >= 0 && index >= rows.blockAfter { + blocked := rows.blocked + ctx := rows.ctx + rows.mu.Unlock() + if blocked != nil { + rows.blockOnce.Do(func() { close(blocked) }) + } + <-ctx.Done() + return ctx.Err() + } + if index >= len(rows.values) { + rows.mu.Unlock() + return io.EOF + } + current := rows.values[index] + rows.index++ + rows.mu.Unlock() + copy(destination, current) + return nil +} + +func (rows *scriptedRows) ColumnTypeDatabaseTypeName(index int) string { + if index < 0 || index >= len(rows.types) { + return "" + } + return rows.types[index] +} + +func (rows *scriptedRows) isClosed() bool { + rows.mu.Lock() + defer rows.mu.Unlock() + return rows.closed +} + +func newScriptedServer(t *testing.T, behavior *scriptedBehavior) *server { + t.Helper() + driverName := fmt.Sprintf("dbx-hive-scripted-%d", scriptedDriverSequence.Add(1)) + sql.Register(driverName, &scriptedDriver{behavior: behavior}) + database, err := sql.Open(driverName, "") + if err != nil { + t.Fatal(err) + } + connection, err := database.Conn(context.Background()) + if err != nil { + database.Close() + t.Fatal(err) + } + server := &server{ + config: connectionConfig{Database: defaultHiveDatabase, ConnectTimeout: time.Second}, + database: database, + connection: connection, + querySessions: map[string]*querySession{}, + } + t.Cleanup(func() { _ = server.disconnect() }) + return server +} + +func rawParams(values map[string]any) map[string]json.RawMessage { + result := make(map[string]json.RawMessage, len(values)) + for key, value := range values { + encoded, err := json.Marshal(value) + if err != nil { + panic(err) + } + result[key] = encoded + } + return result +} + +func TestExecuteQueryUsesHiveServerResultSetSignal(t *testing.T) { + behavior := &scriptedBehavior{} + behavior.query = func(ctx context.Context, query string) (driver.Rows, error) { + switch { + case strings.HasPrefix(query, "WITH source AS"): + return nil, &gohive.NonQueryResult{AffectedRows: 4} + case strings.HasPrefix(query, "SET "): + return newScriptedRows(ctx, []string{"set"}, []string{"STRING"}, [][]driver.Value{{"hive.exec.dynamic.partition=true"}}), nil + default: + return nil, fmt.Errorf("unexpected SQL: %s", query) + } + } + server := newScriptedServer(t, behavior) + + insertResult, err := server.executeQuery(queryOptions{ + SQL: "WITH source AS (SELECT 1) INSERT INTO target SELECT * FROM source", + }) + if err != nil { + t.Fatal(err) + } + if insertResult.AffectedRows != 4 || len(insertResult.Columns) != 0 { + t.Fatalf("unexpected non-query result: %#v", insertResult) + } + + setResult, err := server.executeQuery(queryOptions{SQL: "SET hive.exec.dynamic.partition"}) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(setResult.Rows, [][]any{{"hive.exec.dynamic.partition=true"}}) { + t.Fatalf("unexpected SET result: %#v", setResult.Rows) + } + queries, executions, _, _ := behavior.snapshot() + if len(queries) != 2 || len(executions) != 0 { + t.Fatalf("statements must use HS2 result-set signaling, queries=%v executions=%v", queries, executions) + } +} + +func TestQueryResultsPreserveHiveServerLabelsWithoutDotSplitting(t *testing.T) { + columns := []string{"id", "customer.id", "total + 1", "id"} + columnTypes := []string{"BIGINT", "BIGINT", "DOUBLE", "BIGINT"} + behavior := &scriptedBehavior{} + behavior.query = func(ctx context.Context, _ string) (driver.Rows, error) { + return newScriptedRows(ctx, columns, columnTypes, [][]driver.Value{{int64(1), int64(2), float64(3), int64(4)}, {int64(5), int64(6), float64(7), int64(8)}}), nil + } + server := newScriptedServer(t, behavior) + + ordinary, err := server.executeQuery(queryOptions{SQL: "SELECT * FROM labels"}) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(ordinary.Columns, columns) || !reflect.DeepEqual(ordinary.ColumnTypes, []string{"bigint", "bigint", "double", "bigint"}) { + t.Fatalf("ordinary metadata changed: %#v", ordinary) + } + if len(ordinary.Rows) != 2 || ordinary.Rows[0][0] != "1" || ordinary.Rows[0][3] != "4" { + t.Fatalf("ordinary values changed: %#v", ordinary.Rows) + } + + first, err := server.executeQueryPage(queryOptions{SQL: "SELECT * FROM labels", MaxRows: 2}, 1) + if err != nil { + t.Fatal(err) + } + if first.SessionID == nil || !first.HasMore || !reflect.DeepEqual(first.Columns, columns) { + t.Fatalf("first page metadata changed: %#v", first) + } + second, err := server.fetchQueryPage(*first.SessionID, 1) + if err != nil { + t.Fatal(err) + } + if second.HasMore || second.SessionID != nil || !reflect.DeepEqual(second.Columns, columns) { + t.Fatalf("cached page metadata changed: %#v", second) + } +} + +func TestPagedQueryPreservesDuplicateLeadingValuesAcrossPages(t *testing.T) { + behavior := &scriptedBehavior{} + behavior.query = func(ctx context.Context, _ string) (driver.Rows, error) { + return newScriptedRows( + ctx, + []string{"group_id", "row_id"}, + []string{"BIGINT", "BIGINT"}, + [][]driver.Value{{int64(1), int64(101)}, {int64(1), int64(102)}, {int64(1), int64(103)}}, + ), nil + } + server := newScriptedServer(t, behavior) + + first, err := server.executeQueryPage(queryOptions{SQL: "SELECT group_id, row_id FROM repeated_values", MaxRows: 3}, 2) + if err != nil { + t.Fatal(err) + } + if first.SessionID == nil || !first.HasMore { + t.Fatalf("expected an open cursor after the first page: %#v", first) + } + second, err := server.fetchQueryPage(*first.SessionID, 2) + if err != nil { + t.Fatal(err) + } + allRows := append(append([][]any{}, first.Rows...), second.Rows...) + expected := [][]any{{"1", "101"}, {"1", "102"}, {"1", "103"}} + if !reflect.DeepEqual(allRows, expected) || second.HasMore { + t.Fatalf("duplicate leading values changed across cursor pages: %#v", allRows) + } +} + +func TestPagedQueryTruncatesAndPreservesLegacyJDBCValueSemantics(t *testing.T) { + largeValue := strings.Repeat("x", 256*1024) + createdAt := time.Date(2026, time.August, 11, 10, 11, 12, 345000000, time.UTC) + var sourceRows *scriptedRows + behavior := &scriptedBehavior{} + behavior.query = func(ctx context.Context, _ string) (driver.Rows, error) { + sourceRows = newScriptedRows( + ctx, + []string{"id", "enabled", "payload", "created_at", "complex_value"}, + []string{"BIGINT", "BOOLEAN", "BINARY", "TIMESTAMP", "ARRAY"}, + [][]driver.Value{ + {int64(1), true, []byte{0x00, 0xff}, createdAt, largeValue}, + {int64(2), false, []byte{0x10}, createdAt, "[1,2]"}, + {int64(3), true, []byte{}, createdAt, "map('a',1)"}, + {int64(4), true, []byte{0x01}, createdAt, "extra"}, + }, + ) + return sourceRows, nil + } + server := newScriptedServer(t, behavior) + + first, err := server.executeQueryPage(queryOptions{SQL: "SELECT * FROM values", MaxRows: 3, FetchSize: 2}, 2) + if err != nil { + t.Fatal(err) + } + if !first.HasMore || first.SessionID == nil || len(first.Rows) != 2 { + t.Fatalf("unexpected first page: %#v", first) + } + if !reflect.DeepEqual(first.ColumnTypes, []string{"bigint", "boolean", "binary", "timestamp", "array"}) { + t.Fatalf("unexpected column types: %#v", first.ColumnTypes) + } + if first.Rows[0][0] != "1" || first.Rows[0][1] != "true" || first.Rows[0][2] != "0x00ff" { + t.Fatalf("primitive value types changed: %#v", first.Rows[0]) + } + if first.Rows[0][3] != "2026-08-11 10:11:12.345" || first.Rows[0][4] != largeValue { + t.Fatalf("timestamp or large value changed: %#v", first.Rows[0]) + } + + second, err := server.fetchQueryPage(*first.SessionID, 2) + if err != nil { + t.Fatal(err) + } + if second.HasMore || second.SessionID != nil || !second.Truncated || len(second.Rows) != 1 { + t.Fatalf("unexpected final page: %#v", second) + } + if len(server.querySessions) != 0 || sourceRows == nil || !sourceRows.isClosed() { + t.Fatalf("query session was not closed: sessions=%d rows=%#v", len(server.querySessions), sourceRows) + } +} + +func TestCancelPagedFetchQuarantinesWithoutReplayingSQL(t *testing.T) { + blocked := make(chan struct{}) + behavior := &scriptedBehavior{} + behavior.query = func(ctx context.Context, _ string) (driver.Rows, error) { + rows := newScriptedRows(ctx, []string{"id"}, []string{"BIGINT"}, [][]driver.Value{{int64(1)}, {int64(2)}, {int64(3)}}) + rows.blockAfter = 2 + rows.blocked = blocked + return rows, nil + } + server := newScriptedServer(t, behavior) + first, err := server.executeQueryPage(queryOptions{SQL: "SELECT id FROM slow_table", MaxRows: 10}, 1) + if err != nil { + t.Fatal(err) + } + if first.SessionID == nil { + t.Fatal("expected a paged query session") + } + + result := make(chan error, 1) + go func() { + _, fetchErr := server.fetchQueryPage(*first.SessionID, 1) + result <- fetchErr + }() + select { + case <-blocked: + case <-time.After(2 * time.Second): + t.Fatal("fetch did not reach the blocking row") + } + server.cancelActiveQuery() + fetchErr := <-result + if !errors.Is(fetchErr, context.Canceled) { + t.Fatalf("expected cancellation, got %v", fetchErr) + } + rpcErr := classifyRPCError("fetch_query_page", "session-a", fetchErr) + if rpcErr.Data.Category != "canceled" || rpcErr.Data.SessionDisposition != "quarantine" { + t.Fatalf("unexpected cancellation classification: %#v", rpcErr) + } + queries, _, _, _ := behavior.snapshot() + if len(queries) != 1 { + t.Fatalf("SQL must never be replayed after cancellation: %v", queries) + } + if len(server.querySessions) != 0 { + t.Fatalf("canceled query session was retained: %#v", server.querySessions) + } +} + +func TestPagedFetchHonorsOriginalStatementTimeout(t *testing.T) { + blocked := make(chan struct{}) + behavior := &scriptedBehavior{} + behavior.query = func(ctx context.Context, _ string) (driver.Rows, error) { + rows := newScriptedRows(ctx, []string{"id"}, []string{"BIGINT"}, [][]driver.Value{{int64(1)}, {int64(2)}, {int64(3)}}) + rows.blockAfter = 2 + rows.blocked = blocked + return rows, nil + } + server := newScriptedServer(t, behavior) + first, err := server.executeQueryPage(queryOptions{SQL: "SELECT id FROM slow_table", MaxRows: 10, TimeoutSecs: 1}, 1) + if err != nil { + t.Fatal(err) + } + if first.SessionID == nil { + t.Fatal("expected a paged query session") + } + + started := time.Now() + _, fetchErr := server.fetchQueryPage(*first.SessionID, 1) + if !errors.Is(fetchErr, context.DeadlineExceeded) { + t.Fatalf("expected deadline exceeded, got %v", fetchErr) + } + if elapsed := time.Since(started); elapsed > 3*time.Second { + t.Fatalf("statement timeout was not enforced promptly: %s", elapsed) + } + if len(server.querySessions) != 0 { + t.Fatalf("timed out query session was retained: %#v", server.querySessions) + } +} + +func TestRuntimeSessionsRemainIsolated(t *testing.T) { + behaviorA := &scriptedBehavior{query: func(ctx context.Context, _ string) (driver.Rows, error) { + return newScriptedRows(ctx, []string{"owner"}, []string{"STRING"}, [][]driver.Value{{"a"}}), nil + }} + behaviorB := &scriptedBehavior{query: func(ctx context.Context, _ string) (driver.Rows, error) { + return newScriptedRows(ctx, []string{"owner"}, []string{"STRING"}, [][]driver.Value{{"b"}}), nil + }} + serverA := newScriptedServer(t, behaviorA) + serverB := newScriptedServer(t, behaviorB) + runtimeServer := newRuntimeServer() + runtimeServer.sessions["a"] = &agentSession{server: serverA} + runtimeServer.sessions["b"] = &agentSession{server: serverB} + + resultA, _, err := runtimeServer.dispatch("execute_query", rawParams(map[string]any{ + "agentSessionId": "a", + "sql": "SELECT owner", + })) + if err != nil { + t.Fatal(err) + } + resultB, _, err := runtimeServer.dispatch("execute_query", rawParams(map[string]any{ + "agentSessionId": "b", + "sql": "SELECT owner", + })) + if err != nil { + t.Fatal(err) + } + if resultA.(queryResult).Rows[0][0] != "a" || resultB.(queryResult).Rows[0][0] != "b" { + t.Fatalf("session results crossed: a=%#v b=%#v", resultA, resultB) + } + if err := runtimeServer.closeSession("a"); err != nil { + t.Fatal(err) + } + if _, err := runtimeServer.session("a"); err == nil { + t.Fatal("closed session a is still registered") + } + if _, err := runtimeServer.session("b"); err != nil { + t.Fatalf("closing session a affected session b: %v", err) + } +} + +func TestTransactionFallbackDoesNotReplayFailedStatements(t *testing.T) { + behavior := &scriptedBehavior{} + behavior.exec = func(_ context.Context, query string) (driver.Result, error) { + if strings.Contains(query, "second") { + return nil, io.EOF + } + return driver.RowsAffected(1), nil + } + server := newScriptedServer(t, behavior) + _, err := server.executeStatements(rawParams(map[string]any{ + "statements": []string{"INSERT first", "INSERT second", "INSERT third"}, + }), true) + if !errors.Is(err, io.EOF) { + t.Fatalf("expected connection failure, got %v", err) + } + _, executions, beginCalls, _ := behavior.snapshot() + if !reflect.DeepEqual(executions, []string{"INSERT first", "INSERT second"}) { + t.Fatalf("failed transaction was replayed or continued: %v", executions) + } + if beginCalls == 0 { + t.Fatal("transaction capability was not attempted before fallback") + } + rpcErr := classifyRPCError("execute_transaction", "session-a", err) + if rpcErr.Data.Category != "connection" || rpcErr.Data.SessionDisposition != "quarantine" { + t.Fatalf("unexpected connection failure classification: %#v", rpcErr) + } +} + +func TestTransactionDoesNotFallbackAfterOperationalBeginFailure(t *testing.T) { + behavior := &scriptedBehavior{ + beginErr: errors.New("connection reset while beginning transaction"), + exec: func(context.Context, string) (driver.Result, error) { + return driver.RowsAffected(1), nil + }, + } + server := newScriptedServer(t, behavior) + defer server.disconnect() + params := map[string]json.RawMessage{ + "statements": json.RawMessage(`["INSERT INTO sample VALUES (1)"]`), + } + if _, err := server.executeStatements(params, true); err == nil || !strings.Contains(err.Error(), "connection reset") { + t.Fatalf("expected begin failure, got %v", err) + } + _, executions, _, _ := behavior.snapshot() + if len(executions) != 0 { + t.Fatalf("statements must not execute after begin failure: %#v", executions) + } +} + +func TestExpireIdleQuerySessions(t *testing.T) { + behavior := &scriptedBehavior{query: func(ctx context.Context, _ string) (driver.Rows, error) { + return newScriptedRows(ctx, []string{"id"}, []string{"BIGINT"}, [][]driver.Value{{int64(1)}, {int64(2)}}), nil + }} + server := newScriptedServer(t, behavior) + first, err := server.executeQueryPage(queryOptions{SQL: "SELECT id", MaxRows: 10}, 1) + if err != nil { + t.Fatal(err) + } + if first.SessionID == nil { + t.Fatal("expected a query session") + } + server.querySessions[*first.SessionID].lastAccessed = time.Now().Add(-querySessionIdleTime - time.Second) + if expired := server.expireIdleQuerySessions(time.Now()); expired != 1 { + t.Fatalf("unexpected expired session count: %d", expired) + } + if len(server.querySessions) != 0 { + t.Fatalf("idle session was retained: %#v", server.querySessions) + } +} + +func TestStructuredHiveErrorIncludesServerDiagnostics(t *testing.T) { + err := &gohive.Error{ + Err: errors.New("compile failed"), + Message: "SemanticException", + ErrorCode: 40000, + SQLState: "42000", + } + rpcErr := classifyRPCError("execute_query", "session-a", err) + if rpcErr.Data.Category != "sql" || rpcErr.Data.SQLState != "42000" || rpcErr.Data.VendorCode != 40000 { + t.Fatalf("unexpected Hive diagnostics: %#v", rpcErr) + } + if rpcErr.Data.SessionDisposition != "keep" || rpcErr.Data.OperationOutcome != "unknown" { + t.Fatalf("unexpected SQL failure recovery hints: %#v", rpcErr) + } +} diff --git a/agents/drivers/argo-go/zookeeper_protocol.go b/agents/drivers/argo-go/zookeeper_protocol.go new file mode 100644 index 0000000000..66982d67af --- /dev/null +++ b/agents/drivers/argo-go/zookeeper_protocol.go @@ -0,0 +1,568 @@ +package main + +import ( + "bytes" + "crypto/tls" + "encoding/binary" + "errors" + "fmt" + "io" + "math" + "math/rand/v2" + "net" + "strings" + "sync" + "time" + + "github.com/beltran/gosasl" + "github.com/go-zookeeper/zk" +) + +const ( + zooKeeperProtocolVersion = int32(0) + zooKeeperOpGetData = int32(4) + zooKeeperOpGetChildren2 = int32(12) + zooKeeperOpClose = int32(-11) + zooKeeperOpSetAuth = int32(100) + zooKeeperOpSASL = int32(102) + zooKeeperMaxFrameSize = 16 << 20 + zooKeeperMaxSASLRounds = 8 +) + +var errZooKeeperSessionClosedRequiresSASL = errors.New("ZooKeeper session closed because SASL authentication is required") + +type zooKeeperSASLClient interface { + Start() ([]byte, error) + Step([]byte) ([]byte, error) + Complete() bool + Dispose() +} + +var newZooKeeperSASLClient = func(host string, config connectionConfig) (zooKeeperSASLClient, error) { + service, options := zooKeeperGSSAPIOptions(config) + mechanism, err := gosasl.NewGSSAPIMechanismWithOptions(service, options) + if err != nil { + return nil, err + } + return gosasl.NewSaslClient(host, mechanism), nil +} + +var dialZooKeeperConnection = func(address string, timeout time.Duration, tlsConfig *tls.Config) (net.Conn, error) { + dialer := &net.Dialer{Timeout: timeout} + if tlsConfig == nil { + return dialer.Dial("tcp", address) + } + config := tlsConfig.Clone() + if config.ServerName == "" { + host, _, err := net.SplitHostPort(address) + if err != nil { + return nil, fmt.Errorf("parse ZooKeeper TLS address %q: %w", address, err) + } + config.ServerName = host + } + return tls.DialWithDialer(dialer, "tcp", address, config) +} + +var shuffleZooKeeperServers = func(servers []string) { + rand.Shuffle(len(servers), func(first, second int) { + servers[first], servers[second] = servers[second], servers[first] + }) +} + +func zooKeeperGSSAPIOptions(config connectionConfig) (string, gosasl.GSSAPIOptions) { + service := firstNonEmpty(config.ZooKeeperKerberos.Service, "zookeeper") + options := gssapiOptionsFromKerberos(config.Kerberos) + options.QOP = "auth" + options.AuthorizationID = "" + options.ServiceHost = "" + options.CanonicalizeHost = config.ZooKeeperKerberos.CanonicalHostname + options.ServerName = config.ZooKeeperKerberos.ServerPrincipal + if options.ServerName == "" && config.ZooKeeperKerberos.Realm != "" { + options.ServerName = service + "/_HOST@" + config.ZooKeeperKerberos.Realm + } + return service, options +} + +func connectKerberosZooKeeper( + servers []string, + timeout time.Duration, + tlsConfig *tls.Config, + config connectionConfig, +) (zooKeeperClient, <-chan zk.Event, error) { + if len(servers) == 0 { + return nil, nil, errors.New("ZooKeeper server list is empty") + } + if !config.Kerberos.Enabled { + return nil, nil, errors.New("ZooKeeper Kerberos SASL requires Hive Kerberos credentials") + } + ordered := append([]string(nil), servers...) + shuffleZooKeeperServers(ordered) + var failures []string + for _, address := range ordered { + host, _, err := net.SplitHostPort(address) + if err != nil { + failures = append(failures, fmt.Sprintf("%s: %v", address, err)) + continue + } + connection, err := dialZooKeeperConnection(address, timeout, tlsConfig) + if err != nil { + failures = append(failures, fmt.Sprintf("%s: %v", address, err)) + continue + } + client, err := newProtocolZooKeeperClient(connection, timeout) + if err == nil { + var saslClient zooKeeperSASLClient + saslClient, err = newZooKeeperSASLClient(host, config) + if err == nil { + err = client.authenticateSASL(saslClient) + } + } + if err != nil { + connection.Close() + failures = append(failures, fmt.Sprintf("%s: %v", address, err)) + continue + } + events := make(chan zk.Event, 1) + events <- zk.Event{State: zk.StateHasSession, Server: address} + close(events) + return client, events, nil + } + return nil, nil, fmt.Errorf("connect and authenticate to ZooKeeper: %s", strings.Join(failures, "; ")) +} + +type protocolZooKeeperClient struct { + connection net.Conn + timeout time.Duration + xid int32 + mutex sync.Mutex + closed bool +} + +func newProtocolZooKeeperClient(connection net.Conn, timeout time.Duration) (*protocolZooKeeperClient, error) { + if connection == nil { + return nil, errors.New("ZooKeeper connection is nil") + } + if timeout <= 0 { + timeout = defaultConnectTimeout + } + client := &protocolZooKeeperClient{connection: connection, timeout: timeout} + request := &zooKeeperEncoder{} + request.int32(zooKeeperProtocolVersion) + request.int64(0) + request.int32(zooKeeperTimeoutMillis(timeout)) + request.int64(0) + request.bytes(make([]byte, 16)) + if err := client.writeFrame(request.data()); err != nil { + return nil, fmt.Errorf("send ZooKeeper connect request: %w", err) + } + response, err := client.readFrame() + if err != nil { + return nil, fmt.Errorf("read ZooKeeper connect response: %w", err) + } + decoder := newZooKeeperDecoder(response) + if _, err := decoder.int32(); err != nil { + return nil, fmt.Errorf("decode ZooKeeper protocol version: %w", err) + } + if _, err := decoder.int32(); err != nil { + return nil, fmt.Errorf("decode ZooKeeper session timeout: %w", err) + } + sessionID, err := decoder.int64() + if err != nil { + return nil, fmt.Errorf("decode ZooKeeper session ID: %w", err) + } + if _, err := decoder.bytes(); err != nil { + return nil, fmt.Errorf("decode ZooKeeper session password: %w", err) + } + if sessionID == 0 { + return nil, zk.ErrSessionExpired + } + return client, nil +} + +func zooKeeperTimeoutMillis(timeout time.Duration) int32 { + milliseconds := timeout.Milliseconds() + if milliseconds < 1 { + return 1 + } + if milliseconds > math.MaxInt32 { + return math.MaxInt32 + } + return int32(milliseconds) +} + +func (client *protocolZooKeeperClient) authenticateSASL(saslClient zooKeeperSASLClient) error { + if saslClient == nil { + return errors.New("ZooKeeper SASL client is nil") + } + defer saslClient.Dispose() + token, err := saslClient.Start() + if err != nil { + return fmt.Errorf("start ZooKeeper GSSAPI negotiation: %w", err) + } + for round := 0; round < zooKeeperMaxSASLRounds; round++ { + response, requestErr := client.request(zooKeeperOpSASL, func(encoder *zooKeeperEncoder) { + if token == nil { + encoder.bytes([]byte{}) + return + } + encoder.bytes(token) + }) + if requestErr != nil { + return fmt.Errorf("ZooKeeper SASL round %d: %w", round+1, requestErr) + } + decoder := newZooKeeperDecoder(response) + challenge, decodeErr := decoder.bytes() + if decodeErr != nil { + return fmt.Errorf("decode ZooKeeper SASL round %d: %w", round+1, decodeErr) + } + if saslClient.Complete() { + if len(challenge) != 0 { + return errors.New("ZooKeeper sent an unexpected token after GSSAPI completion") + } + return nil + } + token, err = saslClient.Step(challenge) + if err != nil { + return fmt.Errorf("continue ZooKeeper GSSAPI negotiation at round %d: %w", round+1, err) + } + } + return fmt.Errorf("ZooKeeper GSSAPI negotiation exceeded %d rounds", zooKeeperMaxSASLRounds) +} + +func (client *protocolZooKeeperClient) AddAuth(scheme string, auth []byte) error { + _, err := client.request(zooKeeperOpSetAuth, func(encoder *zooKeeperEncoder) { + encoder.int32(0) + encoder.string(scheme) + encoder.bytes(auth) + }) + return err +} + +func (client *protocolZooKeeperClient) Children(path string) ([]string, *zk.Stat, error) { + response, err := client.request(zooKeeperOpGetChildren2, func(encoder *zooKeeperEncoder) { + encoder.string(path) + encoder.boolean(false) + }) + if err != nil { + return nil, nil, err + } + decoder := newZooKeeperDecoder(response) + children, err := decoder.strings() + if err != nil { + return nil, nil, err + } + stat, err := decoder.stat() + if err != nil { + return nil, nil, err + } + return children, stat, nil +} + +func (client *protocolZooKeeperClient) Get(path string) ([]byte, *zk.Stat, error) { + response, err := client.request(zooKeeperOpGetData, func(encoder *zooKeeperEncoder) { + encoder.string(path) + encoder.boolean(false) + }) + if err != nil { + return nil, nil, err + } + decoder := newZooKeeperDecoder(response) + data, err := decoder.bytes() + if err != nil { + return nil, nil, err + } + stat, err := decoder.stat() + if err != nil { + return nil, nil, err + } + return data, stat, nil +} + +func (client *protocolZooKeeperClient) Close() { + client.mutex.Lock() + defer client.mutex.Unlock() + if client.closed { + return + } + client.closed = true + _ = client.connection.SetDeadline(time.Now().Add(client.timeout)) + client.xid++ + request := &zooKeeperEncoder{} + request.int32(client.xid) + request.int32(zooKeeperOpClose) + _ = client.writeFrame(request.data()) + _ = client.connection.Close() +} + +func (client *protocolZooKeeperClient) request(opcode int32, encodeBody func(*zooKeeperEncoder)) ([]byte, error) { + client.mutex.Lock() + defer client.mutex.Unlock() + if client.closed { + return nil, zk.ErrConnectionClosed + } + client.xid++ + request := &zooKeeperEncoder{} + request.int32(client.xid) + request.int32(opcode) + if encodeBody != nil { + encodeBody(request) + } + if err := client.writeFrame(request.data()); err != nil { + return nil, err + } + response, err := client.readFrame() + if err != nil { + return nil, err + } + decoder := newZooKeeperDecoder(response) + xid, err := decoder.int32() + if err != nil { + return nil, err + } + if xid != client.xid { + return nil, fmt.Errorf("ZooKeeper response XID %d does not match request XID %d", xid, client.xid) + } + if _, err := decoder.int64(); err != nil { + return nil, err + } + code, err := decoder.int32() + if err != nil { + return nil, err + } + if err := zooKeeperError(code); err != nil { + return nil, err + } + return decoder.remaining(), nil +} + +func (client *protocolZooKeeperClient) writeFrame(payload []byte) error { + if len(payload) > zooKeeperMaxFrameSize { + return fmt.Errorf("ZooKeeper request frame is %d bytes, maximum is %d", len(payload), zooKeeperMaxFrameSize) + } + if err := client.connection.SetWriteDeadline(time.Now().Add(client.timeout)); err != nil { + return err + } + header := make([]byte, 4) + binary.BigEndian.PutUint32(header, uint32(len(payload))) + if err := writeAll(client.connection, header); err != nil { + return err + } + return writeAll(client.connection, payload) +} + +func (client *protocolZooKeeperClient) readFrame() ([]byte, error) { + if err := client.connection.SetReadDeadline(time.Now().Add(client.timeout)); err != nil { + return nil, err + } + header := make([]byte, 4) + if _, err := io.ReadFull(client.connection, header); err != nil { + return nil, err + } + length := int(binary.BigEndian.Uint32(header)) + if length < 0 || length > zooKeeperMaxFrameSize { + return nil, fmt.Errorf("ZooKeeper response frame is %d bytes, maximum is %d", length, zooKeeperMaxFrameSize) + } + payload := make([]byte, length) + if _, err := io.ReadFull(client.connection, payload); err != nil { + return nil, err + } + return payload, nil +} + +func writeAll(writer io.Writer, payload []byte) error { + for len(payload) > 0 { + written, err := writer.Write(payload) + if err != nil { + return err + } + if written <= 0 { + return io.ErrShortWrite + } + payload = payload[written:] + } + return nil +} + +func zooKeeperError(code int32) error { + switch code { + case 0: + return nil + case -4: + return zk.ErrConnectionClosed + case -101: + return zk.ErrNoNode + case -102: + return zk.ErrNoAuth + case -112: + return zk.ErrSessionExpired + case -115: + return zk.ErrAuthFailed + case -124: + return errZooKeeperSessionClosedRequiresSASL + default: + return fmt.Errorf("ZooKeeper request failed with error code %d", code) + } +} + +type zooKeeperEncoder struct { + buffer bytes.Buffer +} + +func (encoder *zooKeeperEncoder) int32(value int32) { + var data [4]byte + binary.BigEndian.PutUint32(data[:], uint32(value)) + encoder.buffer.Write(data[:]) +} + +func (encoder *zooKeeperEncoder) int64(value int64) { + var data [8]byte + binary.BigEndian.PutUint64(data[:], uint64(value)) + encoder.buffer.Write(data[:]) +} + +func (encoder *zooKeeperEncoder) boolean(value bool) { + if value { + encoder.buffer.WriteByte(1) + return + } + encoder.buffer.WriteByte(0) +} + +func (encoder *zooKeeperEncoder) string(value string) { + encoder.bytes([]byte(value)) +} + +func (encoder *zooKeeperEncoder) bytes(value []byte) { + if value == nil { + encoder.int32(-1) + return + } + encoder.int32(int32(len(value))) + encoder.buffer.Write(value) +} + +func (encoder *zooKeeperEncoder) data() []byte { + return encoder.buffer.Bytes() +} + +type zooKeeperDecoder struct { + data []byte + offset int +} + +func newZooKeeperDecoder(data []byte) *zooKeeperDecoder { + return &zooKeeperDecoder{data: data} +} + +func (decoder *zooKeeperDecoder) take(length int) ([]byte, error) { + if length < 0 || decoder.offset > len(decoder.data)-length { + return nil, io.ErrUnexpectedEOF + } + value := decoder.data[decoder.offset : decoder.offset+length] + decoder.offset += length + return value, nil +} + +func (decoder *zooKeeperDecoder) int32() (int32, error) { + value, err := decoder.take(4) + if err != nil { + return 0, err + } + return int32(binary.BigEndian.Uint32(value)), nil +} + +func (decoder *zooKeeperDecoder) int64() (int64, error) { + value, err := decoder.take(8) + if err != nil { + return 0, err + } + return int64(binary.BigEndian.Uint64(value)), nil +} + +func (decoder *zooKeeperDecoder) bytes() ([]byte, error) { + length, err := decoder.int32() + if err != nil { + return nil, err + } + if length == -1 { + return nil, nil + } + if length < -1 { + return nil, fmt.Errorf("invalid ZooKeeper buffer length %d", length) + } + value, err := decoder.take(int(length)) + if err != nil { + return nil, err + } + return append([]byte(nil), value...), nil +} + +func (decoder *zooKeeperDecoder) string() (string, error) { + value, err := decoder.bytes() + return string(value), err +} + +func (decoder *zooKeeperDecoder) strings() ([]string, error) { + length, err := decoder.int32() + if err != nil { + return nil, err + } + if length == -1 { + return nil, nil + } + if length < -1 || length > zooKeeperMaxFrameSize/4 { + return nil, fmt.Errorf("invalid ZooKeeper string vector length %d", length) + } + values := make([]string, 0, length) + for index := int32(0); index < length; index++ { + value, valueErr := decoder.string() + if valueErr != nil { + return nil, valueErr + } + values = append(values, value) + } + return values, nil +} + +func (decoder *zooKeeperDecoder) stat() (*zk.Stat, error) { + stat := &zk.Stat{} + var err error + if stat.Czxid, err = decoder.int64(); err != nil { + return nil, err + } + if stat.Mzxid, err = decoder.int64(); err != nil { + return nil, err + } + if stat.Ctime, err = decoder.int64(); err != nil { + return nil, err + } + if stat.Mtime, err = decoder.int64(); err != nil { + return nil, err + } + if stat.Version, err = decoder.int32(); err != nil { + return nil, err + } + if stat.Cversion, err = decoder.int32(); err != nil { + return nil, err + } + if stat.Aversion, err = decoder.int32(); err != nil { + return nil, err + } + if stat.EphemeralOwner, err = decoder.int64(); err != nil { + return nil, err + } + if stat.DataLength, err = decoder.int32(); err != nil { + return nil, err + } + if stat.NumChildren, err = decoder.int32(); err != nil { + return nil, err + } + if stat.Pzxid, err = decoder.int64(); err != nil { + return nil, err + } + return stat, nil +} + +func (decoder *zooKeeperDecoder) remaining() []byte { + return decoder.data[decoder.offset:] +} diff --git a/agents/drivers/argo-go/zookeeper_protocol_test.go b/agents/drivers/argo-go/zookeeper_protocol_test.go new file mode 100644 index 0000000000..e228a0d98c --- /dev/null +++ b/agents/drivers/argo-go/zookeeper_protocol_test.go @@ -0,0 +1,562 @@ +package main + +import ( + "bytes" + "crypto/tls" + "errors" + "fmt" + "io" + "log" + "net" + "os" + "path/filepath" + "testing" + "time" + + "github.com/go-zookeeper/zk" + gsskrb5 "github.com/golang-auth/go-gssapi/v2/krb5" + "github.com/jcmturner/krb5test" +) + +type scriptedZooKeeperSASLClient struct { + complete bool + disposed bool +} + +func (client *scriptedZooKeeperSASLClient) Start() ([]byte, error) { + return []byte("client-initial"), nil +} + +func (client *scriptedZooKeeperSASLClient) Step(challenge []byte) ([]byte, error) { + if string(challenge) != "server-challenge" { + return nil, fmt.Errorf("unexpected challenge %q", challenge) + } + client.complete = true + return []byte("client-final"), nil +} + +func (client *scriptedZooKeeperSASLClient) Complete() bool { return client.complete } +func (client *scriptedZooKeeperSASLClient) Dispose() { client.disposed = true } + +func TestProtocolZooKeeperClientAuthenticatesAndReadsDiscoveryData(t *testing.T) { + clientConnection, serverConnection := net.Pipe() + serverErrors := make(chan error, 1) + go func() { + defer serverConnection.Close() + serverErrors <- serveZooKeeperProtocolTest(serverConnection) + }() + + client, err := newProtocolZooKeeperClient(clientConnection, time.Second) + if err != nil { + t.Fatal(err) + } + saslClient := &scriptedZooKeeperSASLClient{} + if err := client.authenticateSASL(saslClient); err != nil { + t.Fatal(err) + } + if !saslClient.disposed { + t.Fatal("SASL credentials were not disposed") + } + if err := client.AddAuth("digest", []byte("user:password")); err != nil { + t.Fatal(err) + } + children, stat, err := client.Children("/hiveserver2") + if err != nil { + t.Fatal(err) + } + if len(children) != 1 || children[0] != "server-1" || stat.NumChildren != 1 { + t.Fatalf("children=%#v stat=%#v", children, stat) + } + data, stat, err := client.Get("/hiveserver2/server-1") + if err != nil { + t.Fatal(err) + } + if string(data) != "serverUri=hs2.example.com:10000" || stat.DataLength != int32(len(data)) { + t.Fatalf("data=%q stat=%#v", data, stat) + } + client.Close() + if err := <-serverErrors; err != nil { + t.Fatal(err) + } +} + +func TestProtocolZooKeeperClientMapsRequiredSASLError(t *testing.T) { + clientConnection, serverConnection := net.Pipe() + serverErrors := make(chan error, 1) + go func() { + defer serverConnection.Close() + if err := acceptZooKeeperSession(serverConnection); err != nil { + serverErrors <- err + return + } + request, err := readZooKeeperTestFrame(serverConnection) + if err != nil { + serverErrors <- err + return + } + decoder := newZooKeeperDecoder(request) + xid, _ := decoder.int32() + opcode, _ := decoder.int32() + if opcode != zooKeeperOpSASL { + serverErrors <- fmt.Errorf("opcode=%d", opcode) + return + } + serverErrors <- writeZooKeeperTestResponse(serverConnection, xid, -124, nil) + }() + + client, err := newProtocolZooKeeperClient(clientConnection, time.Second) + if err != nil { + t.Fatal(err) + } + saslClient := &scriptedZooKeeperSASLClient{} + err = client.authenticateSASL(saslClient) + client.Close() + if !errors.Is(err, errZooKeeperSessionClosedRequiresSASL) { + t.Fatalf("unexpected error %v", err) + } + if serverErr := <-serverErrors; serverErr != nil { + t.Fatal(serverErr) + } +} + +func TestZooKeeperGSSAPIOptionsDoNotReuseHiveServerIdentity(t *testing.T) { + config := connectionConfig{ + Kerberos: kerberosConfig{ + ServerPrincipal: "hive/_HOST@HIVE.EXAMPLE.COM", + ClientPrincipal: "alice@CLIENT.EXAMPLE.COM", + AuthorizationID: "hive-proxy", + QOP: "auth-conf", + ConfigPath: "/etc/krb5.conf", + CCachePath: "/tmp/alice.ccache", + UseTicketCache: true, + CanonicalHostname: true, + }, + ZooKeeperKerberos: zooKeeperKerberosConfig{ + Enabled: true, + Service: "zookeeper", + Realm: "ZK.EXAMPLE.COM", + CanonicalHostname: false, + }, + } + service, options := zooKeeperGSSAPIOptions(config) + if service != "zookeeper" || options.ServerName != "zookeeper/_HOST@ZK.EXAMPLE.COM" { + t.Fatalf("service=%q options=%#v", service, options) + } + if options.QOP != "auth" || options.AuthorizationID != "" || options.CanonicalizeHost || options.Principal != "alice@CLIENT.EXAMPLE.COM" || options.CCachePath != "/tmp/alice.ccache" { + t.Fatalf("unexpected ZooKeeper GSSAPI options: %#v", options) + } +} + +func TestConnectKerberosZooKeeperFailsOverAndUsesTargetHost(t *testing.T) { + previousDial := dialZooKeeperConnection + previousFactory := newZooKeeperSASLClient + previousShuffle := shuffleZooKeeperServers + t.Cleanup(func() { + dialZooKeeperConnection = previousDial + newZooKeeperSASLClient = previousFactory + shuffleZooKeeperServers = previousShuffle + }) + shuffleZooKeeperServers = func([]string) {} + + clientConnection, serverConnection := net.Pipe() + serverErrors := make(chan error, 1) + go func() { + defer serverConnection.Close() + if err := acceptZooKeeperSession(serverConnection); err != nil { + serverErrors <- err + return + } + if err := expectZooKeeperSASLRound(serverConnection, "client-initial", []byte("server-challenge")); err != nil { + serverErrors <- err + return + } + if err := expectZooKeeperSASLRound(serverConnection, "client-final", nil); err != nil { + serverErrors <- err + return + } + request, err := readZooKeeperTestFrame(serverConnection) + if err != nil { + serverErrors <- err + return + } + decoder := newZooKeeperDecoder(request) + _, _ = decoder.int32() + opcode, _ := decoder.int32() + if opcode != zooKeeperOpClose { + serverErrors <- fmt.Errorf("unexpected close opcode %d", opcode) + return + } + serverErrors <- nil + }() + + var dialed []string + dialZooKeeperConnection = func(address string, _ time.Duration, _ *tls.Config) (net.Conn, error) { + dialed = append(dialed, address) + if address == "first.example.com:2181" { + return nil, errors.New("first unavailable") + } + return clientConnection, nil + } + var targetHost string + newZooKeeperSASLClient = func(host string, _ connectionConfig) (zooKeeperSASLClient, error) { + targetHost = host + return &scriptedZooKeeperSASLClient{}, nil + } + + config := connectionConfig{ + Kerberos: kerberosConfig{Enabled: true}, + ZooKeeperKerberos: zooKeeperKerberosConfig{Enabled: true}, + } + client, events, err := connectKerberosZooKeeper( + []string{"first.example.com:2181", "second.example.com:2181"}, + time.Second, + nil, + config, + ) + if err != nil { + t.Fatal(err) + } + if len(dialed) != 2 || dialed[0] != "first.example.com:2181" || dialed[1] != "second.example.com:2181" { + t.Fatalf("dialed = %#v", dialed) + } + if targetHost != "second.example.com" { + t.Fatalf("target host = %q", targetHost) + } + event := <-events + if event.State != zk.StateHasSession || event.Server != "second.example.com:2181" { + t.Fatalf("event = %#v", event) + } + client.Close() + if err := <-serverErrors; err != nil { + t.Fatal(err) + } +} + +func TestZooKeeperKerberosSASLWithMiniKDC(t *testing.T) { + logger := log.New(io.Discard, "", 0) + kdc, err := krb5test.NewKDC(map[string][]string{ + "alice": nil, + "zookeeper/localhost": nil, + }, logger) + if err != nil { + t.Fatal(err) + } + kdc.KRB5Conf.LibDefaults.UDPPreferenceLimit = 1 + kdc.Start() + defer kdc.Close() + + directory := t.TempDir() + configPath := filepath.Join(directory, "krb5.conf") + keytabPath := filepath.Join(directory, "zookeeper.keytab") + configContents := fmt.Sprintf(`[libdefaults] + default_realm = %s + dns_lookup_realm = false + dns_lookup_kdc = false + rdns = false + udp_preference_limit = 1 + default_tgs_enctypes = aes256-cts-hmac-sha1-96 + default_tkt_enctypes = aes256-cts-hmac-sha1-96 + permitted_enctypes = aes256-cts-hmac-sha1-96 + +[realms] + %s = { + kdc = %s + } +`, kdc.Realm, kdc.Realm, kdc.TCPListener.Addr().String()) + if err := os.WriteFile(configPath, []byte(configContents), 0o600); err != nil { + t.Fatal(err) + } + keytabBytes, err := kdc.Keytab.Marshal() + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(keytabPath, keytabBytes, 0o600); err != nil { + t.Fatal(err) + } + t.Setenv("KRB5_KTNAME", keytabPath) + t.Setenv("KRB5_CLIENT_KTNAME", keytabPath) + + clientConnection, serverConnection := net.Pipe() + serverErrors := make(chan error, 1) + servicePrincipal := "zookeeper/localhost@" + kdc.Realm + go func() { + defer serverConnection.Close() + serverErrors <- serveKerberosZooKeeperSASL(serverConnection, servicePrincipal) + }() + + protocolClient, err := newProtocolZooKeeperClient(clientConnection, 5*time.Second) + if err != nil { + t.Fatal(err) + } + config := connectionConfig{ + Kerberos: kerberosConfig{ + Enabled: true, + ClientPrincipal: "alice@" + kdc.Realm, + ConfigPath: configPath, + Password: kdc.Principals["alice"].Password, + DisablePAFXFAST: true, + }, + ZooKeeperKerberos: zooKeeperKerberosConfig{ + Enabled: true, + Service: "zookeeper", + ServerPrincipal: servicePrincipal, + CanonicalHostname: false, + }, + } + saslClient, err := newZooKeeperSASLClient("localhost", config) + if err != nil { + t.Fatal(err) + } + if err := protocolClient.authenticateSASL(saslClient); err != nil { + t.Fatal(err) + } + protocolClient.Close() + if err := <-serverErrors; err != nil { + t.Fatal(err) + } +} + +func serveZooKeeperProtocolTest(connection net.Conn) error { + if err := acceptZooKeeperSession(connection); err != nil { + return err + } + if err := expectZooKeeperSASLRound(connection, "client-initial", []byte("server-challenge")); err != nil { + return err + } + if err := expectZooKeeperSASLRound(connection, "client-final", nil); err != nil { + return err + } + request, err := readZooKeeperTestFrame(connection) + if err != nil { + return err + } + decoder := newZooKeeperDecoder(request) + xid, _ := decoder.int32() + opcode, _ := decoder.int32() + authType, _ := decoder.int32() + scheme, _ := decoder.string() + auth, _ := decoder.bytes() + if opcode != zooKeeperOpSetAuth || authType != 0 || scheme != "digest" || string(auth) != "user:password" { + return fmt.Errorf("unexpected auth request opcode=%d type=%d scheme=%q auth=%q", opcode, authType, scheme, auth) + } + if err := writeZooKeeperTestResponse(connection, xid, 0, nil); err != nil { + return err + } + request, err = readZooKeeperTestFrame(connection) + if err != nil { + return err + } + decoder = newZooKeeperDecoder(request) + xid, _ = decoder.int32() + opcode, _ = decoder.int32() + path, _ := decoder.string() + watch, _ := decoder.take(1) + if opcode != zooKeeperOpGetChildren2 || path != "/hiveserver2" || !bytes.Equal(watch, []byte{0}) { + return fmt.Errorf("unexpected children request opcode=%d path=%q watch=%v", opcode, path, watch) + } + body := &zooKeeperEncoder{} + body.int32(1) + body.string("server-1") + encodeZooKeeperTestStat(body, 0, 1) + if err := writeZooKeeperTestResponse(connection, xid, 0, body.data()); err != nil { + return err + } + request, err = readZooKeeperTestFrame(connection) + if err != nil { + return err + } + decoder = newZooKeeperDecoder(request) + xid, _ = decoder.int32() + opcode, _ = decoder.int32() + path, _ = decoder.string() + watch, _ = decoder.take(1) + if opcode != zooKeeperOpGetData || path != "/hiveserver2/server-1" || !bytes.Equal(watch, []byte{0}) { + return fmt.Errorf("unexpected get request opcode=%d path=%q watch=%v", opcode, path, watch) + } + data := []byte("serverUri=hs2.example.com:10000") + body = &zooKeeperEncoder{} + body.bytes(data) + encodeZooKeeperTestStat(body, int32(len(data)), 0) + if err := writeZooKeeperTestResponse(connection, xid, 0, body.data()); err != nil { + return err + } + request, err = readZooKeeperTestFrame(connection) + if err != nil { + return err + } + decoder = newZooKeeperDecoder(request) + _, _ = decoder.int32() + opcode, _ = decoder.int32() + if opcode != zooKeeperOpClose { + return fmt.Errorf("unexpected close opcode %d", opcode) + } + return nil +} + +func serveKerberosZooKeeperSASL(connection net.Conn, servicePrincipal string) error { + if err := acceptZooKeeperSession(connection); err != nil { + return err + } + acceptor := gsskrb5.NewKrb5Mech() + if err := acceptor.Accept(servicePrincipal); err != nil { + return err + } + xid, token, err := readZooKeeperSASLRequest(connection) + if err != nil { + return err + } + apReply, err := acceptor.Continue(token) + if err != nil { + return err + } + if !acceptor.IsEstablished() { + return errors.New("Kerberos acceptor context was not established") + } + if err := writeZooKeeperSASLResponse(connection, xid, apReply); err != nil { + return err + } + xid, token, err = readZooKeeperSASLRequest(connection) + if err != nil { + return err + } + if len(token) != 0 { + return fmt.Errorf("expected empty post-AP-REP token, got %x", token) + } + securityChallenge, err := acceptor.Wrap([]byte{1, 0, 0, 0}, false) + if err != nil { + return err + } + if err := writeZooKeeperSASLResponse(connection, xid, securityChallenge); err != nil { + return err + } + xid, token, err = readZooKeeperSASLRequest(connection) + if err != nil { + return err + } + securityResponse, sealed, err := acceptor.Unwrap(token) + if err != nil { + return err + } + if sealed || len(securityResponse) < 4 || !bytes.Equal(securityResponse[:4], []byte{1, 0, 0, 0}) { + return fmt.Errorf("invalid GSSAPI security-layer response sealed=%v payload=%x", sealed, securityResponse) + } + if len(securityResponse[4:]) == 0 { + return errors.New("GSSAPI authorization identity is empty") + } + if err := writeZooKeeperSASLResponse(connection, xid, nil); err != nil { + return err + } + request, err := readZooKeeperTestFrame(connection) + if err != nil { + return err + } + decoder := newZooKeeperDecoder(request) + _, _ = decoder.int32() + opcode, _ := decoder.int32() + if opcode != zooKeeperOpClose { + return fmt.Errorf("unexpected close opcode %d", opcode) + } + return nil +} + +func readZooKeeperSASLRequest(connection net.Conn) (int32, []byte, error) { + request, err := readZooKeeperTestFrame(connection) + if err != nil { + return 0, nil, err + } + decoder := newZooKeeperDecoder(request) + xid, err := decoder.int32() + if err != nil { + return 0, nil, err + } + opcode, err := decoder.int32() + if err != nil { + return 0, nil, err + } + if opcode != zooKeeperOpSASL { + return 0, nil, fmt.Errorf("unexpected SASL opcode %d", opcode) + } + token, err := decoder.bytes() + return xid, token, err +} + +func writeZooKeeperSASLResponse(connection net.Conn, xid int32, token []byte) error { + body := &zooKeeperEncoder{} + body.bytes(token) + return writeZooKeeperTestResponse(connection, xid, 0, body.data()) +} + +func acceptZooKeeperSession(connection net.Conn) error { + request, err := readZooKeeperTestFrame(connection) + if err != nil { + return err + } + decoder := newZooKeeperDecoder(request) + version, _ := decoder.int32() + _, _ = decoder.int64() + timeout, _ := decoder.int32() + sessionID, _ := decoder.int64() + password, _ := decoder.bytes() + if version != 0 || timeout <= 0 || sessionID != 0 || len(password) != 16 { + return fmt.Errorf("invalid connect request version=%d timeout=%d session=%d password=%d", version, timeout, sessionID, len(password)) + } + response := &zooKeeperEncoder{} + response.int32(0) + response.int32(timeout) + response.int64(42) + response.bytes(make([]byte, 16)) + return writeZooKeeperTestFrame(connection, response.data()) +} + +func expectZooKeeperSASLRound(connection net.Conn, wantToken string, responseToken []byte) error { + request, err := readZooKeeperTestFrame(connection) + if err != nil { + return err + } + decoder := newZooKeeperDecoder(request) + xid, _ := decoder.int32() + opcode, _ := decoder.int32() + token, _ := decoder.bytes() + if opcode != zooKeeperOpSASL || string(token) != wantToken { + return fmt.Errorf("unexpected SASL request opcode=%d token=%q", opcode, token) + } + body := &zooKeeperEncoder{} + body.bytes(responseToken) + return writeZooKeeperTestResponse(connection, xid, 0, body.data()) +} + +func encodeZooKeeperTestStat(encoder *zooKeeperEncoder, dataLength, children int32) { + encoder.int64(1) + encoder.int64(2) + encoder.int64(3) + encoder.int64(4) + encoder.int32(5) + encoder.int32(6) + encoder.int32(7) + encoder.int64(8) + encoder.int32(dataLength) + encoder.int32(children) + encoder.int64(9) +} + +func writeZooKeeperTestResponse(connection net.Conn, xid, code int32, body []byte) error { + response := &zooKeeperEncoder{} + response.int32(xid) + response.int64(0) + response.int32(code) + response.buffer.Write(body) + return writeZooKeeperTestFrame(connection, response.data()) +} + +func readZooKeeperTestFrame(connection net.Conn) ([]byte, error) { + client := &protocolZooKeeperClient{connection: connection, timeout: time.Second} + return client.readFrame() +} + +func writeZooKeeperTestFrame(connection net.Conn, payload []byte) error { + client := &protocolZooKeeperClient{connection: connection, timeout: time.Second} + return client.writeFrame(payload) +} + +var _ zooKeeperClient = (*protocolZooKeeperClient)(nil) +var _ = zk.StateSaslAuthenticated diff --git a/agents/drivers/argo-go/zookeeper_tls.go b/agents/drivers/argo-go/zookeeper_tls.go new file mode 100644 index 0000000000..4ad6265a45 --- /dev/null +++ b/agents/drivers/argo-go/zookeeper_tls.go @@ -0,0 +1,240 @@ +package main + +import ( + "bytes" + "crypto/tls" + "crypto/x509" + "encoding/pem" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + + keystore "github.com/pavlo-v-chernykh/keystore-go/v4" + pkcs12 "software.sslmate.com/src/go-pkcs12" +) + +func buildZooKeeperTLSConfig(values map[string]string) (*tls.Config, error) { + if !parameterBool(values, "zookeepersslenable") { + return nil, nil + } + config := &tls.Config{ + MinVersion: tls.VersionTLS12, + ServerName: parameter(values, "zookeeperservername"), + } + trustStoreLocation := parameter(values, "zookeepertruststorelocation") + if trustStoreLocation != "" { + certificates, err := loadTrustStore( + trustStoreLocation, + parameter(values, "zookeepertruststorepassword"), + parameter(values, "zookeepertruststoretype"), + ) + if err != nil { + return nil, fmt.Errorf("load ZooKeeper truststore: %w", err) + } + pool := x509.NewCertPool() + for _, certificate := range certificates { + pool.AddCert(certificate) + } + config.RootCAs = pool + } + keyStoreLocation := parameter(values, "zookeeperkeystorelocation") + if keyStoreLocation != "" { + certificate, err := loadClientKeyStore( + keyStoreLocation, + parameter(values, "zookeeperkeystorepassword"), + parameter(values, "zookeeperkeystoretype"), + ) + if err != nil { + return nil, fmt.Errorf("load ZooKeeper keystore: %w", err) + } + config.Certificates = []tls.Certificate{certificate} + } + if parameterBool(values, "zookeepersslinsecureskipverify") { + config.InsecureSkipVerify = true + } + return config, nil +} + +func loadTrustStore(path, password, storeType string) ([]*x509.Certificate, error) { + contents, err := os.ReadFile(path) + if err != nil { + return nil, err + } + switch normalizedStoreType(storeType, path) { + case "PEM": + return parsePEMCertificates(contents) + case "PKCS12": + certificates, err := pkcs12.DecodeTrustStore(contents, password) + if err == nil { + return certificates, nil + } + _, certificate, chain, chainErr := pkcs12.DecodeChain(contents, password) + if chainErr != nil { + return nil, err + } + return append([]*x509.Certificate{certificate}, chain...), nil + case "JKS": + store, err := loadJKS(contents, password) + if err != nil { + return nil, err + } + var certificates []*x509.Certificate + for _, alias := range store.Aliases() { + switch { + case store.IsTrustedCertificateEntry(alias): + entry, getErr := store.GetTrustedCertificateEntry(alias) + if getErr != nil { + return nil, getErr + } + certificate, parseErr := x509.ParseCertificate(entry.Certificate.Content) + if parseErr != nil { + return nil, parseErr + } + certificates = append(certificates, certificate) + case store.IsPrivateKeyEntry(alias): + chain, getErr := store.GetPrivateKeyEntryCertificateChain(alias) + if getErr != nil { + return nil, getErr + } + for _, entry := range chain { + certificate, parseErr := x509.ParseCertificate(entry.Content) + if parseErr != nil { + return nil, parseErr + } + certificates = append(certificates, certificate) + } + } + } + if len(certificates) == 0 { + return nil, errors.New("JKS truststore contains no certificates") + } + return certificates, nil + default: + return nil, fmt.Errorf("unsupported store type %q", storeType) + } +} + +func loadClientKeyStore(path, password, storeType string) (tls.Certificate, error) { + contents, err := os.ReadFile(path) + if err != nil { + return tls.Certificate{}, err + } + switch normalizedStoreType(storeType, path) { + case "PEM": + return tls.X509KeyPair(contents, contents) + case "PKCS12": + privateKey, certificate, chain, err := pkcs12.DecodeChain(contents, password) + if err != nil { + return tls.Certificate{}, err + } + result := tls.Certificate{PrivateKey: privateKey, Leaf: certificate} + result.Certificate = append(result.Certificate, certificate.Raw) + for _, entry := range chain { + result.Certificate = append(result.Certificate, entry.Raw) + } + return result, nil + case "JKS": + store, err := loadJKS(contents, password) + if err != nil { + return tls.Certificate{}, err + } + passwordBytes := []byte(password) + defer clear(passwordBytes) + for _, alias := range store.Aliases() { + if !store.IsPrivateKeyEntry(alias) { + continue + } + entry, getErr := store.GetPrivateKeyEntry(alias, passwordBytes) + if getErr != nil { + return tls.Certificate{}, getErr + } + privateKey, parseErr := parsePrivateKey(entry.PrivateKey) + if parseErr != nil { + return tls.Certificate{}, parseErr + } + result := tls.Certificate{PrivateKey: privateKey} + for index, certificate := range entry.CertificateChain { + result.Certificate = append(result.Certificate, certificate.Content) + if index == 0 { + result.Leaf, _ = x509.ParseCertificate(certificate.Content) + } + } + if len(result.Certificate) == 0 { + return tls.Certificate{}, errors.New("JKS private key entry has no certificate chain") + } + return result, nil + } + return tls.Certificate{}, errors.New("JKS keystore contains no private key entry") + default: + return tls.Certificate{}, fmt.Errorf("unsupported store type %q", storeType) + } +} + +func normalizedStoreType(storeType, path string) string { + value := strings.ToUpper(strings.TrimSpace(storeType)) + switch value { + case "P12", "PFX", "PKCS#12": + return "PKCS12" + case "X509", "X.509": + return "PEM" + case "": + switch strings.ToLower(filepath.Ext(path)) { + case ".jks": + return "JKS" + case ".p12", ".pfx", ".pkcs12": + return "PKCS12" + default: + return "PEM" + } + default: + return value + } +} + +func loadJKS(contents []byte, password string) (keystore.KeyStore, error) { + store := keystore.New() + passwordBytes := []byte(password) + defer clear(passwordBytes) + if err := store.Load(bytes.NewReader(contents), passwordBytes); err != nil { + return keystore.KeyStore{}, err + } + return store, nil +} + +func parsePEMCertificates(contents []byte) ([]*x509.Certificate, error) { + var certificates []*x509.Certificate + for len(contents) > 0 { + block, rest := pem.Decode(contents) + if block == nil { + break + } + contents = rest + if block.Type != "CERTIFICATE" { + continue + } + certificate, err := x509.ParseCertificate(block.Bytes) + if err != nil { + return nil, err + } + certificates = append(certificates, certificate) + } + if len(certificates) == 0 { + return nil, errors.New("PEM truststore contains no certificates") + } + return certificates, nil +} + +func parsePrivateKey(contents []byte) (any, error) { + if value, err := x509.ParsePKCS8PrivateKey(contents); err == nil { + return value, nil + } + if value, err := x509.ParsePKCS1PrivateKey(contents); err == nil { + return value, nil + } + if value, err := x509.ParseECPrivateKey(contents); err == nil { + return value, nil + } + return nil, errors.New("unsupported private key encoding") +} diff --git a/agents/drivers/argo-go/zookeeper_tls_test.go b/agents/drivers/argo-go/zookeeper_tls_test.go new file mode 100644 index 0000000000..07e65d002f --- /dev/null +++ b/agents/drivers/argo-go/zookeeper_tls_test.go @@ -0,0 +1,162 @@ +package main + +import ( + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "os" + "path/filepath" + "testing" + "time" + + keystore "github.com/pavlo-v-chernykh/keystore-go/v4" + pkcs12 "software.sslmate.com/src/go-pkcs12" +) + +func testZooKeeperCertificate(t *testing.T) (*rsa.PrivateKey, *x509.Certificate, []byte) { + t.Helper() + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatal(err) + } + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "zk.example.com"}, + DNSNames: []string{"zk.example.com"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment | x509.KeyUsageCertSign, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth, x509.ExtKeyUsageClientAuth}, + IsCA: true, + } + raw, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey) + if err != nil { + t.Fatal(err) + } + certificate, err := x509.ParseCertificate(raw) + if err != nil { + t.Fatal(err) + } + return privateKey, certificate, raw +} + +func TestBuildZooKeeperTLSConfigFromPEM(t *testing.T) { + _, _, raw := testZooKeeperCertificate(t) + path := filepath.Join(t.TempDir(), "trust.pem") + if err := os.WriteFile(path, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: raw}), 0o600); err != nil { + t.Fatal(err) + } + config, err := buildZooKeeperTLSConfig(map[string]string{ + "zookeepersslenable": "true", + "zookeepertruststorelocation": path, + "zookeepertruststoretype": "PEM", + "zookeeperservername": "zk.example.com", + }) + if err != nil { + t.Fatal(err) + } + if config == nil || config.RootCAs == nil || config.ServerName != "zk.example.com" { + t.Fatalf("unexpected TLS config: %#v", config) + } +} + +func TestBuildZooKeeperTLSConfigFromJKS(t *testing.T) { + privateKey, certificate, raw := testZooKeeperCertificate(t) + password := []byte("changeit") + store := keystore.New() + if err := store.SetTrustedCertificateEntry("ca", keystore.TrustedCertificateEntry{ + CreationTime: time.Now(), + Certificate: keystore.Certificate{Type: "X509", Content: raw}, + }); err != nil { + t.Fatal(err) + } + encodedKey, err := x509.MarshalPKCS8PrivateKey(privateKey) + if err != nil { + t.Fatal(err) + } + if err := store.SetPrivateKeyEntry("client", keystore.PrivateKeyEntry{ + CreationTime: time.Now(), + PrivateKey: encodedKey, + CertificateChain: []keystore.Certificate{ + {Type: "X509", Content: certificate.Raw}, + }, + }, password); err != nil { + t.Fatal(err) + } + path := filepath.Join(t.TempDir(), "client.jks") + file, err := os.Create(path) + if err != nil { + t.Fatal(err) + } + if err := store.Store(file, password); err != nil { + file.Close() + t.Fatal(err) + } + if err := file.Close(); err != nil { + t.Fatal(err) + } + config, err := buildZooKeeperTLSConfig(map[string]string{ + "zookeepersslenable": "true", + "zookeepertruststorelocation": path, + "zookeepertruststorepassword": string(password), + "zookeepertruststoretype": "JKS", + "zookeeperkeystorelocation": path, + "zookeeperkeystorepassword": string(password), + "zookeeperkeystoretype": "JKS", + }) + if err != nil { + t.Fatal(err) + } + if config.RootCAs == nil || len(config.Certificates) != 1 || config.Certificates[0].PrivateKey == nil { + t.Fatalf("unexpected JKS TLS config: %#v", config) + } +} + +func TestBuildZooKeeperTLSConfigFromPKCS12(t *testing.T) { + privateKey, certificate, _ := testZooKeeperCertificate(t) + password := "changeit" + keyStore, err := pkcs12.Modern.Encode(privateKey, certificate, nil, password) + if err != nil { + t.Fatal(err) + } + trustStore, err := pkcs12.Modern.EncodeTrustStore([]*x509.Certificate{certificate}, password) + if err != nil { + t.Fatal(err) + } + directory := t.TempDir() + keyPath := filepath.Join(directory, "client.p12") + trustPath := filepath.Join(directory, "trust.p12") + if err := os.WriteFile(keyPath, keyStore, 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(trustPath, trustStore, 0o600); err != nil { + t.Fatal(err) + } + config, err := buildZooKeeperTLSConfig(map[string]string{ + "zookeepersslenable": "true", + "zookeepertruststorelocation": trustPath, + "zookeepertruststorepassword": password, + "zookeepertruststoretype": "PKCS12", + "zookeeperkeystorelocation": keyPath, + "zookeeperkeystorepassword": password, + "zookeeperkeystoretype": "PKCS12", + }) + if err != nil { + t.Fatal(err) + } + if config.RootCAs == nil || len(config.Certificates) != 1 { + t.Fatalf("unexpected PKCS12 TLS config: %#v", config) + } +} + +func TestZooKeeperTLSRequiresExistingStore(t *testing.T) { + if _, err := buildZooKeeperTLSConfig(map[string]string{ + "zookeepersslenable": "true", + "zookeepertruststorelocation": filepath.Join(t.TempDir(), "missing.jks"), + }); err == nil { + t.Fatal("expected missing truststore error") + } +} diff --git a/agents/go-common/gohive/hive.go b/agents/go-common/gohive/hive.go index 52f0170bfd..ca82c53e9c 100644 --- a/agents/go-common/gohive/hive.go +++ b/agents/go-common/gohive/hive.go @@ -856,13 +856,11 @@ func (c *cursor) executeSync(ctx context.Context, query string) { return } if !success(safeStatus(responseExecute.GetStatus())) { - status := safeStatus(responseExecute.GetStatus()) - c.Err = &Error{ - Err: errors.New("Error while executing query: " + status.String()), - Message: status.GetErrorMessage(), - ErrorCode: int(status.GetErrorCode()), - SQLState: status.GetSqlState(), - } + // status.String() dumps raw Thrift struct fields; pointer fields (SqlState, + // ErrorMessage) render as hex addresses like 0x2c45f4a70e10. Route through + // hiveStatusError so failures surface the server's real diagnostics + // (message, SQLState, error code) the way every other call site already does. + c.Err = hiveStatusError("executing statement", responseExecute.GetStatus()) return } diff --git a/agents/metadata-constraint-coverage.tsv b/agents/metadata-constraint-coverage.tsv index 543491ad23..41246c025a 100644 --- a/agents/metadata-constraint-coverage.tsv +++ b/agents/metadata-constraint-coverage.tsv @@ -1,5 +1,6 @@ driver strategy scope reason access intentional-fallback java-jdbc-metadata Uses Access JDBC metadata; no portable stable server-side paging API, common constraints filter locally. +argo-go shared-fallback native-go Uses HiveServer2 metadata APIs and the ArgoDB routine catalog views, then applies stable filtering and paging in the native agent with a SHOW TABLES fallback. bigquery native-pushdown java-sql Uses dataset INFORMATION_SCHEMA.TABLES with type, normal-name filter, stable order, and literal LIMIT/OFFSET. cassandra-go shared-fallback native-go Uses cassandra-gocql-driver schema metadata, then applies stable filtering and paging in the native agent. dameng native-pushdown java-sql Uses ALL_OBJECTS with type, filter, stable order, and LIMIT/OFFSET with legacy fallback. diff --git a/agents/scripts/validate_agents.py b/agents/scripts/validate_agents.py index cff62f80fe..3a8e6c7f64 100644 --- a/agents/scripts/validate_agents.py +++ b/agents/scripts/validate_agents.py @@ -16,6 +16,7 @@ "cassandra": "drivers/cassandra-go", "duckdb": "drivers/duckdb", "hive": "drivers/hive-go", + "argo": "drivers/argo-go", "oracle": "drivers/oracle-go", "kingbase": "drivers/kingbase-go", "iotdb": "drivers/iotdb", diff --git a/agents/scripts/version_agent_artifacts.py b/agents/scripts/version_agent_artifacts.py index 1a53a13721..038f5bb6fd 100644 --- a/agents/scripts/version_agent_artifacts.py +++ b/agents/scripts/version_agent_artifacts.py @@ -4,7 +4,7 @@ from pathlib import Path -NATIVE_DRIVERS = ("cassandra", "hive", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "duckdb", "rabbitmq", "rocketmq", "zookeeper", "tdengine", "etcd", "etcd2", "sqlite-worker") +NATIVE_DRIVERS = ("cassandra", "hive", "argo", "oracle", "xugu", "kingbase", "iotdb", "neo4j", "vastbase", "duckdb", "rabbitmq", "rocketmq", "zookeeper", "tdengine", "etcd", "etcd2", "sqlite-worker") PLATFORMS = ( "macos-aarch64", "macos-x64", diff --git a/agents/versions.json b/agents/versions.json index bc49d8f4f6..ac84fb59f0 100644 --- a/agents/versions.json +++ b/agents/versions.json @@ -1,5 +1,5 @@ { - "access": "0.1.49", +"access": "0.1.49", "bigquery": "0.1.53", "cassandra": "0.1.46", "dameng": "0.1.59", @@ -15,6 +15,7 @@ "h2": "0.1.52", "h2-legacy": "0.1.22", "hive": "0.1.53", + "argo": "0.1.0", "ignite": "0.1.5", "ignite3": "0.1.5", "spark": "0.1.24", diff --git a/apps/desktop/src/lib/__tests__/sql/sqlStatementRanges.spec.ts b/apps/desktop/src/lib/__tests__/sql/sqlStatementRanges.spec.ts index 10eb926fb9..150303585f 100644 --- a/apps/desktop/src/lib/__tests__/sql/sqlStatementRanges.spec.ts +++ b/apps/desktop/src/lib/__tests__/sql/sqlStatementRanges.spec.ts @@ -109,6 +109,34 @@ BEGIN NULL; END;`; +// Mirrors the SP_ETL_LOG procedure from the 2026-09-01 ArgoDB session: PL/SQL-style +// body with semicolons inside INSERT statements plus a block comment header. The +// frontend splitter must keep the whole definition as ONE statement — otherwise +// "execute current statement" sends only the first fragment +// (`CREATE ... IS BEGIN INSERT INTO ...`) and the server returns 42000 + 1101. +const argoProcedureFixture = `CREATE OR REPLACE PROCEDURE SP_ETL_LOG +( + II_DATDATE IN INT, --数据日期 + IV_SCHEMA_NAME IN STRING --模式名 +) +/**************************************** +@AUTHOR:xiangxu +#0.20150906-xiangxu-处理执行信息插入日志表 +*****************************************/ +IS +BEGIN + INSERT INTO dws.ETL_LOG + ( + DATA_DATE, + SCHEMA_NAME + ) + VALUES + ( + II_DATDATE, + IV_SCHEMA_NAME + ); +END;`; + const gaussDbDollarQuotedFunctionScript = `DROP FUNCTION IF EXISTS dbx_issue_4572_tmp_md5_uuid; CREATE OR REPLACE FUNCTION dbx_issue_4572_tmp_md5_uuid (v_str IN TEXT) RETURNS varchar(36) LANGUAGE PLPGSQL IMMUTABLE AS $function$ @@ -444,6 +472,17 @@ describe("splitSqlStatementRanges", () => { expect(rangeSqlTexts(splitSqlStatementRanges(oracleIssue2405PlSql, "oracle"))).toEqual([oracleIssue2405PlSql]); }); + it("keeps ArgoDB PL/SQL procedure bodies together (frontend splitter, mirrors backend argo_split tests)", () => { + expect(rangeSqlTexts(splitSqlStatementRanges(argoProcedureFixture, "argo"))).toEqual([argoProcedureFixture]); + expect(hasMultipleExecutionTargets(argoProcedureFixture, "argo")).toBe(false); + expect(rangeSqlTexts(executableStatementRanges(argoProcedureFixture, "argo"))).toEqual([argoProcedureFixture]); + }); + + it("statement at cursor inside an ArgoDB procedure body returns the whole definition", () => { + const range = statementRangeAtCursor(argoProcedureFixture, indexOf(argoProcedureFixture, "ETL_LOG", 2), "argo"); + expect(range?.sql.trim()).toBe(argoProcedureFixture.trim()); + }); + it("keeps consecutive nested Oracle blocks inside their procedure", () => { expect(rangeSqlTexts(splitSqlStatementRanges(`${oracleConsecutiveNestedBlocks}\nSELECT 1;`, "oracle"))).toEqual([oracleConsecutiveNestedBlocks, "SELECT 1"]); }); diff --git a/apps/desktop/src/lib/connection/agentDriverInstallHint.ts b/apps/desktop/src/lib/connection/agentDriverInstallHint.ts index 5316a58ca1..9850049804 100644 --- a/apps/desktop/src/lib/connection/agentDriverInstallHint.ts +++ b/apps/desktop/src/lib/connection/agentDriverInstallHint.ts @@ -33,7 +33,8 @@ export function hasInstalledAgentVersion(drivers: readonly AgentDriverInstallSta export function agentDriverInstallKey(dbType: DatabaseType | undefined, driverProfile?: string, context?: AgentDriverInstallContext): string | undefined { if (dbType === "sqlite") return context?.ssh ? "sqlite-worker" : undefined; - if (dbType === "kyuubi" || dbType === "impala" || dbType === "argo") return "hive"; + // argo owns its dedicated argo-go agent — only kyuubi/impala still share hive-go. + if (dbType === "kyuubi" || dbType === "impala") return "hive"; if (dbType === "oracle") return "oracle"; if (dbType === "h2") return "h2"; if (dbType === "mongodb") return "mongodb"; diff --git a/apps/desktop/src/lib/sql/sqlStatementRanges.ts b/apps/desktop/src/lib/sql/sqlStatementRanges.ts index 5df972a4fa..ed497315d5 100644 --- a/apps/desktop/src/lib/sql/sqlStatementRanges.ts +++ b/apps/desktop/src/lib/sql/sqlStatementRanges.ts @@ -266,7 +266,11 @@ const ALTER_BODY_KEYWORDS = new Set(["ADD", "ALTER", "COMMENT", "DROP", "MODIFY" const CLICKHOUSE_ALTER_TABLE_HEADER = /^ALTER\s+TABLE\s+(?:(?:[A-Za-z_][\w$]*|`(?:``|[^`])+`|"(?:""|[^"])+")\s*\.\s*)?(?:[A-Za-z_][\w$]*|`(?:``|[^`])+`|"(?:""|[^"])+")(?:\s+ON\s+CLUSTER\s+(?:[A-Za-z_][\w$]*|`(?:``|[^`])+`|"(?:""|[^"])+"|'(?:''|[^'])+'))?\s*$/i; const SET_OPERATION_KEYWORDS = new Set(["UNION", "INTERSECT", "EXCEPT", "MINUS"]); const SET_OPERATION_MODIFIER_KEYWORDS = new Set(["ALL", "DISTINCT"]); -const ORACLE_LIKE_PL_SQL_DATABASES: ReadonlySet = new Set(["oracle", "dameng", "gaussdb", "yashandb", "oscar", "oceanbase-oracle", "xugu"]); +// Mirrors the backend list in dbx-core/src/sql.rs is_oracle_like_database — keep both +// in sync. ArgoDB (Transwarp Hive/Inceptor fork) ships a PL/SQL-compatible procedure +// language (`CREATE [OR REPLACE] PROCEDURE ... IS BEGIN ... END;`), so its statement +// ranges must stay whole instead of splitting at every body semicolon. +const ORACLE_LIKE_PL_SQL_DATABASES: ReadonlySet = new Set(["oracle", "dameng", "gaussdb", "yashandb", "oscar", "oceanbase-oracle", "xugu", "argo"]); const MYSQL_ROUTINE_BLOCK_DATABASES: ReadonlySet = new Set(["mysql", "doris", "starrocks", "manticoresearch", "goldendb"]); const MYSQL_CREATE_TABLE_OPTION_DATABASES: ReadonlySet = new Set(["mysql", "doris", "starrocks", "manticoresearch", "goldendb", "gbase"]); const MYSQL_ROUTINE_OBJECT_TYPES = new Set(["PROCEDURE", "FUNCTION", "TRIGGER", "EVENT"]); diff --git a/crates/dbx-core/assets/database-drivers.manifest.json b/crates/dbx-core/assets/database-drivers.manifest.json index 2097a6620f..288ef8602d 100644 --- a/crates/dbx-core/assets/database-drivers.manifest.json +++ b/crates/dbx-core/assets/database-drivers.manifest.json @@ -1844,12 +1844,12 @@ }, { "dbType": "argo", - "label": "ArgoDB (Transwarp)", + "label": "星环Argo", "runtimeMode": "agent", "mcpMode": "bridge", - "agentKey": "hive", - "driverStoreVisible": false, - "driverStoreOrder": 23, + "agentKey": "argo", + "driverStoreVisible": true, + "driverStoreOrder": 50, "singleConnectionPool": false, "metadataConnectionScoped": false, "skipTcpProbe": true, diff --git a/crates/dbx-core/src/sql.rs b/crates/dbx-core/src/sql.rs index ef261cd96e..1b81b2a4fa 100644 --- a/crates/dbx-core/src/sql.rs +++ b/crates/dbx-core/src/sql.rs @@ -225,6 +225,11 @@ impl SqlDialectProfile { | DatabaseType::Oscar | DatabaseType::OceanbaseOracle | DatabaseType::Xugu + // ArgoDB (Transwarp fork of Hive/Inceptor) ships a PL/SQL-compatible procedure + // language: `CREATE [OR REPLACE] PROCEDURE ... IS BEGIN ... END;`. Treat it like + // Oracle for statement splitting so semicolons inside the procedure body are not + // misinterpreted as client-side statement terminators. + | DatabaseType::Argo ) } } @@ -3623,6 +3628,149 @@ END;"; assert_eq!(split_sql_statements_for_database(sql, DatabaseType::Xugu), vec![sql.to_string()]); } + #[test] + fn argo_split_keeps_create_procedure_body_together() { + // ArgoDB (Transwarp) accepts PL/SQL-style procedure definitions whose body contains + // semicolons. The statement splitter must keep those semicolons inside the body rather + // than emitting one fragment per statement — otherwise the agent sends only the first + // `CREATE ... IS BEGIN INSERT INTO ...` fragment and the server responds with + // 42000 + vendorCode 1101 (see screenshots from the 2026-09-01 session). + let sql = "\ +CREATE OR REPLACE PROCEDURE SP_ETL_LOG +( + II_DATDATE IN INT, + IV_SCHEMA_NAME IN STRING, + IV_PROCEDURE_NAME IN STRING, + II_STEP_ID IN INT, + IV_STEP_DESC IN STRING, + II_STEP_FLAG IN INT, + II_START_TIME IN TIMESTAMP +) +/**************************************** +@AUTHOR:xiangxu +@CREATE-DATE:2015-09-06 +@DESCRIPTION:处理执行信息插入日志表(ETL_LOG) +@MODIFICATION HISTORY: +#0.20150906-xiangxu-处理执行信息插入日志表 +*****************************************/ +IS +BEGIN + INSERT INTO dws.ETL_LOG + ( + DATA_DATE, + SCHEMA_NAME, + PROCEDURE_NAME, + STEP_ID, + STEP_DESC, + STEP_FLAG, + START_TIME + ) + VALUES + ( + II_DATDATE, + IV_SCHEMA_NAME, + IV_PROCEDURE_NAME, + II_STEP_ID, + IV_STEP_DESC, + II_STEP_FLAG, + II_START_TIME + ); +END;"; + + assert_eq!(split_sql_statements_for_database(sql, DatabaseType::Argo), vec![sql.to_string()]); + } + + #[test] + fn argo_split_keeps_procedure_with_dash_line_comments() { + // Mirrors the SP_ETL_LOG definition in the 2026-09-01 screenshot where each + // parameter row carries a trailing `--中文` line comment. The PL/SQL tokenizer must + // still recognise the procedure as a block; otherwise the splitter emits the + // parameter list, "COMMIT", and "END" as three separate statements (the failure + // mode shown in the 21:22:59 screenshot). + let sql = "\ +CREATE OR REPLACE PROCEDURE SP_ETL_LOG +( + II_DATDATE IN INT, --数据日期 + IV_SCHEMA_NAME IN STRING, --模式名 + IV_PROCEDURE_NAME IN STRING, --存储过程名称 + II_STEP_ID IN INT, --任务号 + IV_STEP_DESC IN STRING, --任务描述 + II_STEP_FLAG IN INT, --执行状态 + II_START_TIME IN TIMESTAMP --起始时间 +); +END;"; + + let statements = split_sql_statements_for_database(sql, DatabaseType::Argo); + assert_eq!(statements, vec![sql.to_string()], "splitter produced: {statements:#?}"); + } + + #[test] + fn argo_split_handles_full_user_procedure_with_block_comment_and_at_meta() { + // Mirrors the EXACT procedure text shown in the 2026-09-01 21:22:59 screenshot: + // - parameter list with trailing `--中文` comments + // - block comment `/********...********/` containing `@AUTHOR:` and `#0.20150906-...` + // which are NOT standard SQL comment markers (ArgoDB/Hive ignores them, but the + // splitter must still treat them as part of the surrounding block comment) + // - `IS BEGIN INSERT INTO ... ; END;` + // The result panel showed three split fragments; if this test passes, the splitter + // is fine and the issue is on the DBX.app side. + let sql = "CREATE OR REPLACE PROCEDURE SP_ETL_LOG\n\ +(\n\ + II_DATDATE IN INT, --数据日期\n\ + IV_SCHEMA_NAME IN STRING, --模式名\n\ + IV_PROCEDURE_NAME IN STRING, --存储过程名称\n\ + II_STEP_ID IN INT, --任务号\n\ + IV_STEP_DESC IN STRING, --任务描述\n\ + II_STEP_FLAG IN INT, --执行状态\n\ + II_START_TIME IN TIMESTAMP --起始时间\n\ +)\n\ +/****************************************\n\ +@AUTHOR:xiangxu\n\ +@CREATE-DATE:2015-09-06\n\ +@DESCRIPTION:处理执行信息插入日志表(ETL_LOG)\n\ +@MODIFICATION HISTORY:\n\ +#0.20150906-xiangxu-处理执行信息插入日志表\n\ +*****************************************/\n\ +IS\n\ +BEGIN INSERT INTO dws.ETL_LOG\n\ + (\n\ + DATA_DATE, --数据日期\n\ + SCHEMA_NAME, --模式名\n\ + PROCEDURE_NAME, --存储过程名称\n\ + STEP_ID, --任务号\n\ + STEP_DESC, --任务描述\n\ + STEP_FLAG, --执行状态\n\ + START_TIME --起始时间\n\ + )\n\ + ;\n\ +END;"; + + let statements = split_sql_statements_for_database(sql, DatabaseType::Argo); + assert_eq!(statements.len(), 1, "splitter produced {statements:#?}"); + assert_eq!(statements[0], sql); + } + + #[test] + fn argo_split_handles_minimal_user_procedure_no_body_semicolon() { + // Mirrors the user's 2026-09-02 10:09 procedure: simple CREATE OR REPLACE PROCEDURE + // body with no semicolon after INSERT INTO ... SELECT, only after the SELECT expression. + let sql = "CREATE OR REPLACE PROCEDURE SP_TEST_PART_A\n\ + (\n\ + II_DATADATE IN INT\n\ + )\n\ + IS\n\ + I_DATADATE INT;\n\ + BEGIN\n\ + I_DATADATE := II_DATADATE;\n\ + INSERT INTO dws.TEST_PART_PROC\n\ + SELECT 1, 'testA1', I_DATADATE;\n\ + END;"; + + let statements = split_sql_statements_for_database(sql, DatabaseType::Argo); + assert_eq!(statements.len(), 1, "splitter produced {statements:#?}"); + assert_eq!(statements[0], sql); + } + #[test] fn gaussdb_split_ignores_psql_controls_before_anonymous_block_per_issue_6468() { let block = "\ diff --git a/crates/dbx-core/src/sql_dialect/tests.rs b/crates/dbx-core/src/sql_dialect/tests.rs index 6bd7266b88..448dc2821c 100644 --- a/crates/dbx-core/src/sql_dialect/tests.rs +++ b/crates/dbx-core/src/sql_dialect/tests.rs @@ -115,6 +115,11 @@ fn qualifies_schema_only_for_schema_aware_databases() { "\"DBX_TEST\".\"PRODUCTS\"" ); assert_eq!(qualified_table_name(Some(DatabaseType::Oscar), Some("SYSDBA"), "EMPLOYEE"), "\"SYSDBA\".\"EMPLOYEE\""); + // ArgoDB (Transwarp fork of Hive) shares Hive's backtick identifier syntax; it must be + // classified as schema-aware AND quoted with backticks (not the default `"..."`, which + // ArgoDB parses as a string literal — see dbx-argo-double-quote-bug memory note). + assert_eq!(qualified_table_name(Some(DatabaseType::Argo), Some("dws"), "etl_log"), "`dws`.`etl_log`"); + assert_eq!(qualified_table_name(Some(DatabaseType::Argo), None, "etl_log"), "`etl_log`"); assert_eq!(qualified_table_name(Some(DatabaseType::Informix), Some("xtdpcky"), "users"), "xtdpcky.users"); assert_eq!(qualified_table_name(Some(DatabaseType::Sqlite), Some("analytics"), "users"), "\"analytics\".\"users\""); assert_eq!(qualified_table_name(Some(DatabaseType::Jdbc), Some("cbsdw_dwd"), "dwd_test_df"), "dwd_test_df"); diff --git a/crates/dbx-core/tests/database_capabilities.rs b/crates/dbx-core/tests/database_capabilities.rs index a84855d7db..45b05c190a 100644 --- a/crates/dbx-core/tests/database_capabilities.rs +++ b/crates/dbx-core/tests/database_capabilities.rs @@ -101,7 +101,7 @@ fn maps_agent_database_types_to_driver_keys() { assert_eq!(agent_key(&DatabaseType::Hive, None), Some("hive")); assert_eq!(agent_key(&DatabaseType::Kyuubi, None), Some("hive")); assert_eq!(agent_key(&DatabaseType::Impala, None), Some("hive")); - assert_eq!(agent_key(&DatabaseType::Argo, None), Some("hive")); + assert_eq!(agent_key(&DatabaseType::Argo, None), Some("argo")); assert_eq!(agent_key(&DatabaseType::Tdengine, None), Some("tdengine")); assert_eq!(agent_key(&DatabaseType::Iotdb, None), Some("iotdb")); assert_eq!(agent_key(&DatabaseType::Yashandb, None), Some("yashandb")); diff --git a/plugins/connection-types/argo.yaml b/plugins/connection-types/argo.yaml index 9e06340488..fa51a04060 100644 --- a/plugins/connection-types/argo.yaml +++ b/plugins/connection-types/argo.yaml @@ -2,12 +2,12 @@ schemaVersion: 1 order: 515 dbType: argo rustVariant: Argo -label: ArgoDB (Transwarp) +label: 星环Argo runtimeMode: agent mcpMode: bridge -agentKey: hive -driverStoreVisible: false -driverStoreOrder: 23 +agentKey: argo +driverStoreVisible: true +driverStoreOrder: 50 singleConnectionPool: false metadataConnectionScoped: false skipTcpProbe: true