Add skill enable controls and data factory SQL skill
This commit is contained in:
@@ -0,0 +1,188 @@
|
||||
"""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)
|
||||
Reference in New Issue
Block a user