189 lines
6.7 KiB
Python
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)
|