Files
zk-data-agent/skills/data-factory-sql/kyuubi_client.py
T

189 lines
6.7 KiB
Python

"""Kyuubi HTTP API client for Xiaomi Data Factory.
API spec: https://mi.feishu.cn/wiki/Svthwh7g9isyKbkIXO0czfX5nef
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
from typing import Any, Iterator
import requests
CLUSTER_BASE_URLS = {
"cnbj1": "http://proxy-service-http-cnbj1-dp.api.xiaomi.net",
"cnbj2": "http://proxy-service-http-cnbj2-dp.api.xiaomi.net",
"alsgp0": "http://proxy-service-http-alisgp0-dp.api.xiaomi.net",
"ksyru0": "http://proxy-service-http-ksyru0-dp.api.xiaomi.net",
"azamsprc0": "http://proxy-service-http-azamsprc0-dp.api.xiaomi.net",
}
# State values per official doc (NOT the demo's QUEUED/FAILED/CANCELLED).
TERMINAL_OK = {"FINISHED"}
TERMINAL_FAIL = {"ERROR", "TIMEOUT", "CLOSED"}
NON_TERMINAL = {"PENDING", "RUNNING"}
class KyuubiError(Exception):
"""Base error for Kyuubi client."""
class AuthError(KyuubiError):
"""Token rejected or missing permission."""
class QueryError(KyuubiError):
"""Server-side SQL execution error."""
@dataclass
class QueryResult:
columns: list[dict[str, str]] = field(default_factory=list)
rows: list[list[Any]] = field(default_factory=list)
query_id: str | None = None
engine: str | None = None
elapsed_ms: int | None = None
class KyuubiClient:
def __init__(
self,
token: str,
base_url: str = CLUSTER_BASE_URLS["cnbj1"],
engine: str = "auto",
catalog: str | None = None,
schema: str | None = None,
timeout: int = 30,
poll_interval: float = 2.0,
query_timeout_seconds: int = 600,
):
self.base_url = base_url.rstrip("/")
self.token = token
self.engine = engine
self.timeout = timeout
self.poll_interval = poll_interval
self.query_timeout_seconds = query_timeout_seconds
self.session = requests.Session()
headers = {
"X-SqlProxy-User": token,
"X-SqlProxy-Engine": engine,
"Content-Type": "text/plain;charset=utf-8",
}
if catalog:
headers["X-SqlProxy-Catalog"] = catalog
if schema:
headers["X-SqlProxy-Schema"] = schema
self.session.headers.update(headers)
def _check(self, payload: dict[str, Any]) -> dict[str, Any]:
meta = payload.get("meta") or {}
code = meta.get("errCode", 0)
if code != 0:
msg = meta.get("errMsg", "")
if code == 4007402:
raise AuthError(f"auth failed (errCode {code}): {msg}")
raise QueryError(f"errCode {code}: {msg}")
return payload.get("data") or {}
def submit(self, sql: str) -> tuple[str, str | None]:
"""POST /query → (queryId, engine_picked)."""
url = f"{self.base_url}/olap/api/v2/statement/query"
resp = self.session.post(url, data=sql.encode("utf-8"), timeout=self.timeout)
if resp.status_code == 401:
raise AuthError("HTTP 401: token rejected")
if resp.status_code != 200:
raise QueryError(f"submit HTTP {resp.status_code}: {resp.text[:2000]}")
data = self._check(resp.json())
qid = data.get("queryId")
if not qid:
raise QueryError(f"no queryId returned: {data}")
return qid, data.get("engine")
def poll_until_done(self, query_id: str) -> tuple[str, dict[str, Any]]:
"""Poll getStatusAndLog until terminal state. Returns (latest_query_id, status_data).
Important: every status response returns a NEW nextQueryId; we must use the
latest one for the next poll AND for the eventual fetchResult call.
"""
url = f"{self.base_url}/olap/api/v2/statement/getStatusAndLog"
current = query_id
deadline = time.time() + self.query_timeout_seconds
last_data: dict[str, Any] = {}
while True:
resp = self.session.post(
url, params={"queryId": current}, timeout=self.timeout
)
if resp.status_code != 200:
raise QueryError(f"status HTTP {resp.status_code}: {resp.text[:2000]}")
data = self._check(resp.json())
last_data = data
state = data.get("state", "")
next_id = data.get("nextQueryId") or current
if state in TERMINAL_OK:
return next_id, data
if state in TERMINAL_FAIL:
err = data.get("exceptionMsg") or data.get("simpleExceptionMsg") or state
raise QueryError(f"query {state}: {err}")
if state not in NON_TERMINAL:
raise QueryError(f"unknown state '{state}': {data}")
if time.time() > deadline:
self.close(current)
raise QueryError(
f"query exceeded {self.query_timeout_seconds}s, last state={state}"
)
current = next_id
time.sleep(self.poll_interval)
def fetch_results(self, query_id: str) -> Iterator[dict[str, Any]]:
"""Iterate fetchResult chunks. Yields {columns, rows, state}."""
url = f"{self.base_url}/olap/api/v2/statement/fetchResult"
current: str | None = query_id
while current:
resp = self.session.post(
url, params={"queryId": current}, timeout=self.timeout
)
if resp.status_code != 200:
raise QueryError(f"fetch HTTP {resp.status_code}: {resp.text[:2000]}")
data = self._check(resp.json())
yield data
nxt = data.get("nextResultQueryId") or ""
current = nxt if nxt else None
def close(self, query_id: str) -> None:
url = f"{self.base_url}/olap/api/v2/statement/close"
try:
self.session.post(
url, params={"queryId": query_id}, timeout=self.timeout
)
except requests.RequestException:
pass # best-effort
def execute(self, sql: str) -> QueryResult:
"""Submit → wait → fetch all → return aggregated QueryResult."""
t0 = time.time()
qid, engine_picked = self.submit(sql)
try:
final_qid, _status = self.poll_until_done(qid)
cols: list[dict[str, str]] = []
rows: list[list[Any]] = []
for chunk in self.fetch_results(final_qid):
if not cols and chunk.get("columns"):
cols = chunk["columns"]
rows.extend(chunk.get("rows") or [])
return QueryResult(
columns=cols,
rows=rows,
query_id=final_qid,
engine=engine_picked,
elapsed_ms=int((time.time() - t0) * 1000),
)
finally:
self.close(qid)