"""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)