arrowbricks 0.2.0__tar.gz → 0.3.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/PKG-INFO +1 -1
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/pyproject.toml +1 -1
- arrowbricks-0.3.0/scripts/benchmark_client_reuse.py +126 -0
- arrowbricks-0.3.0/scripts/benchmark_simulated.py +157 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/src/arrowbricks/client.py +81 -34
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/tests/conftest.py +5 -2
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/tests/test_client.py +81 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/uv.lock +1 -1
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/.github/workflows/ci.yml +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/.github/workflows/release.yml +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/.gitignore +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/AGENTS.md +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/LICENSE +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/README.md +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/examples/azure_auth.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/examples/basic.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/examples/cursor_paging.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/examples/fastapi_sse.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/examples/fastapi_sse_pivot.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/prek.toml +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/src/arrowbricks/__init__.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/src/arrowbricks/_streaming.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/src/arrowbricks/cursor.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/tests/test_cursor.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/tests/test_result_set.py +0 -0
- {arrowbricks-0.2.0 → arrowbricks-0.3.0}/tests/test_streaming.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: arrowbricks
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.3.0
|
|
4
4
|
Summary: Runs SQL against a Databricks SQL warehouse via the Statement Execution API and hands you the result as Arrow -- a DB-API-ish Cursor (fetchone/fetchmany/fetchall/fetchall_arrow) or NDJSON streaming. Single Arrow engine (arro3), no DuckDB.
|
|
5
5
|
Project-URL: Repository, https://github.com/bmsuisse/arrowbricks
|
|
6
6
|
License-Expression: MIT
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "arrowbricks"
|
|
3
|
-
version = "0.
|
|
3
|
+
version = "0.3.0"
|
|
4
4
|
description = "Runs SQL against a Databricks SQL warehouse via the Statement Execution API and hands you the result as Arrow -- a DB-API-ish Cursor (fetchone/fetchmany/fetchall/fetchall_arrow) or NDJSON streaming. Single Arrow engine (arro3), no DuckDB."
|
|
5
5
|
readme = "README.md"
|
|
6
6
|
license = "MIT"
|
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
"""Before/after benchmark for the persistent-client + warehouse-check-cache
|
|
2
|
+
optimizations (see client.py's _get_http_client/_ensure_warehouse_running).
|
|
3
|
+
|
|
4
|
+
Runs the same small query N times two ways against a REAL Databricks
|
|
5
|
+
warehouse:
|
|
6
|
+
- "cold": a fresh DatabricksClient per query (simulates the old
|
|
7
|
+
per-call-httpx.AsyncClient() + per-call warehouse-status-GET behavior).
|
|
8
|
+
- "warm": one DatabricksClient reused across all N queries (the new
|
|
9
|
+
default behavior).
|
|
10
|
+
|
|
11
|
+
Usage:
|
|
12
|
+
export DATABRICKS_HOST=... # or falls back to `databricks auth env`
|
|
13
|
+
export DATABRICKS_WAREHOUSE_ID=...
|
|
14
|
+
export DATABRICKS_TOKEN=... # a personal access token
|
|
15
|
+
uv run python scripts/benchmark_client_reuse.py [--n 10] [--sql "SELECT 1"]
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import argparse
|
|
21
|
+
import asyncio
|
|
22
|
+
import os
|
|
23
|
+
import shutil
|
|
24
|
+
import statistics
|
|
25
|
+
import subprocess
|
|
26
|
+
import sys
|
|
27
|
+
import time
|
|
28
|
+
|
|
29
|
+
sys.path.insert(0, "src")
|
|
30
|
+
|
|
31
|
+
from arrowbricks import DatabricksClient # noqa: E402
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def _databricks_auth_env() -> dict[str, str]:
|
|
35
|
+
"""Falls back to `databricks auth env` (the CLI's own resolved
|
|
36
|
+
credentials) if DATABRICKS_HOST/TOKEN aren't set directly -- lets this
|
|
37
|
+
script work whether the user set env vars or ran `databricks configure`."""
|
|
38
|
+
if os.environ.get("DATABRICKS_HOST") and (
|
|
39
|
+
os.environ.get("DATABRICKS_TOKEN") or os.environ.get("DATABRICKS_CLIENT_ID")
|
|
40
|
+
):
|
|
41
|
+
return dict(os.environ)
|
|
42
|
+
databricks_cli = shutil.which("databricks")
|
|
43
|
+
if not databricks_cli:
|
|
44
|
+
return dict(os.environ)
|
|
45
|
+
try:
|
|
46
|
+
out = subprocess.run( # noqa: S603 -- fixed args, no untrusted input
|
|
47
|
+
[databricks_cli, "auth", "env"], capture_output=True, text=True, check=True
|
|
48
|
+
).stdout
|
|
49
|
+
except Exception:
|
|
50
|
+
return dict(os.environ)
|
|
51
|
+
env = dict(os.environ)
|
|
52
|
+
for line in out.splitlines():
|
|
53
|
+
line = line.strip()
|
|
54
|
+
if line.startswith("export "):
|
|
55
|
+
line = line[len("export ") :]
|
|
56
|
+
if "=" in line:
|
|
57
|
+
k, _, v = line.partition("=")
|
|
58
|
+
env.setdefault(k.strip(), v.strip().strip('"'))
|
|
59
|
+
return env
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
async def _run_n_cold(host: str, warehouse_id: str, token: str, sql: str, n: int) -> list[float]:
|
|
63
|
+
"""Old behavior: a fresh client (and fresh http connection pool + a
|
|
64
|
+
fresh warehouse-status check) per query."""
|
|
65
|
+
times = []
|
|
66
|
+
for _ in range(n):
|
|
67
|
+
client = DatabricksClient(host, warehouse_id, token=token)
|
|
68
|
+
start = time.perf_counter()
|
|
69
|
+
await client.execute_json_statement(sql)
|
|
70
|
+
times.append(time.perf_counter() - start)
|
|
71
|
+
await client.aclose()
|
|
72
|
+
return times
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
async def _run_n_warm(host: str, warehouse_id: str, token: str, sql: str, n: int) -> list[float]:
|
|
76
|
+
"""New behavior: one client, one connection pool, warehouse status
|
|
77
|
+
checked once then cached for the TTL window."""
|
|
78
|
+
times = []
|
|
79
|
+
async with DatabricksClient(host, warehouse_id, token=token) as client:
|
|
80
|
+
for _ in range(n):
|
|
81
|
+
start = time.perf_counter()
|
|
82
|
+
await client.execute_json_statement(sql)
|
|
83
|
+
times.append(time.perf_counter() - start)
|
|
84
|
+
return times
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def _report(label: str, times: list[float]) -> None:
|
|
88
|
+
print(f"{label}: n={len(times)} mean={statistics.mean(times):.3f}s median={statistics.median(times):.3f}s "
|
|
89
|
+
f"min={min(times):.3f}s max={max(times):.3f}s total={sum(times):.3f}s")
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
async def main() -> None:
|
|
93
|
+
parser = argparse.ArgumentParser()
|
|
94
|
+
parser.add_argument("--n", type=int, default=10)
|
|
95
|
+
parser.add_argument("--sql", default="SELECT 1")
|
|
96
|
+
args = parser.parse_args()
|
|
97
|
+
|
|
98
|
+
env = _databricks_auth_env()
|
|
99
|
+
host = env.get("DATABRICKS_HOST", "adb-8956277663194228.8.azuredatabricks.net")
|
|
100
|
+
warehouse_id = env.get("DATABRICKS_WAREHOUSE_ID", "c397040753b46093")
|
|
101
|
+
token = env.get("DATABRICKS_TOKEN")
|
|
102
|
+
if not token:
|
|
103
|
+
raise SystemExit(
|
|
104
|
+
"No DATABRICKS_TOKEN found (env var or `databricks auth env`). "
|
|
105
|
+
"Set one and re-run -- this script needs a real, working credential."
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
print(f"host={host} warehouse_id={warehouse_id} n={args.n} sql={args.sql!r}\n")
|
|
109
|
+
|
|
110
|
+
# Warm the warehouse first (cold-start time shouldn't pollute either measurement).
|
|
111
|
+
warmup = DatabricksClient(host, warehouse_id, token=token)
|
|
112
|
+
await warmup.execute_json_statement(args.sql)
|
|
113
|
+
await warmup.aclose()
|
|
114
|
+
|
|
115
|
+
cold = await _run_n_cold(host, warehouse_id, token, args.sql, args.n)
|
|
116
|
+
_report("cold (fresh client per query)", cold)
|
|
117
|
+
|
|
118
|
+
warm = await _run_n_warm(host, warehouse_id, token, args.sql, args.n)
|
|
119
|
+
_report("warm (one reused client) ", warm)
|
|
120
|
+
|
|
121
|
+
speedup = statistics.mean(cold) / statistics.mean(warm) if statistics.mean(warm) > 0 else float("inf")
|
|
122
|
+
print(f"\nmean speedup: {speedup:.2f}x")
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
if __name__ == "__main__":
|
|
126
|
+
asyncio.run(main())
|
|
@@ -0,0 +1,157 @@
|
|
|
1
|
+
"""SIMULATED before/after benchmark for the persistent-client +
|
|
2
|
+
warehouse-check-cache optimizations (client.py's _get_http_client /
|
|
3
|
+
_ensure_warehouse_running) -- runs entirely against a respx-mocked
|
|
4
|
+
transport, no real Databricks warehouse or credentials needed.
|
|
5
|
+
|
|
6
|
+
This is NOT a measurement against a real network -- respx intercepts at the
|
|
7
|
+
transport layer, so there's no actual socket, TCP handshake, or TLS
|
|
8
|
+
negotiation happening. To make the comparison meaningful anyway, this script
|
|
9
|
+
injects realistic artificial latency for exactly the two things the
|
|
10
|
+
optimizations remove:
|
|
11
|
+
|
|
12
|
+
- `_HANDSHAKE_COST_S`: paid once per NEW httpx.AsyncClient (simulates the
|
|
13
|
+
TCP+TLS handshake a brand-new connection to the Databricks host pays).
|
|
14
|
+
"cold" mode creates a fresh DatabricksClient (and thus a fresh
|
|
15
|
+
httpx.AsyncClient) per query, so it pays this every time. "warm" mode
|
|
16
|
+
reuses one DatabricksClient across all queries, paying it exactly once.
|
|
17
|
+
- `_WAREHOUSE_CHECK_COST_S`: the warehouse-status GET's round-trip time.
|
|
18
|
+
This one isn't injected by this script at all -- it falls out of the
|
|
19
|
+
REAL client.py code: "cold" mode's fresh client has no cached
|
|
20
|
+
confirmation, so _ensure_warehouse_running does the GET every time;
|
|
21
|
+
"warm" mode's cache (see warehouse_confirmed_running_ttl_s) skips it
|
|
22
|
+
after the first call. The respx route itself sleeps
|
|
23
|
+
_WAREHOUSE_CHECK_COST_S before responding, same for both modes -- the
|
|
24
|
+
difference in total time comes only from how many times each mode hits
|
|
25
|
+
that route.
|
|
26
|
+
|
|
27
|
+
Every other mocked endpoint (statement submission, chunk resolve/fetch) has
|
|
28
|
+
zero added delay in both modes, since that cost doesn't differ between them
|
|
29
|
+
and would only dilute the comparison.
|
|
30
|
+
|
|
31
|
+
Usage:
|
|
32
|
+
uv run python scripts/benchmark_simulated.py [--n 20]
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
from __future__ import annotations
|
|
36
|
+
|
|
37
|
+
import argparse
|
|
38
|
+
import asyncio
|
|
39
|
+
import statistics
|
|
40
|
+
import sys
|
|
41
|
+
import time
|
|
42
|
+
|
|
43
|
+
sys.path.insert(0, "src")
|
|
44
|
+
sys.path.insert(0, "tests")
|
|
45
|
+
|
|
46
|
+
import httpx # noqa: E402
|
|
47
|
+
import respx # noqa: E402
|
|
48
|
+
from conftest import HOST, WAREHOUSE_ID, build_chunk_bytes # noqa: E402
|
|
49
|
+
|
|
50
|
+
from arrowbricks import DatabricksClient # noqa: E402
|
|
51
|
+
|
|
52
|
+
_HANDSHAKE_COST_S = 0.06 # a fresh HTTPS connection's TCP+TLS setup, typical for a cross-region call
|
|
53
|
+
_WAREHOUSE_CHECK_COST_S = 0.04 # one small API round-trip
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def _install_routes(router: respx.Router, statement_id: str) -> None:
|
|
57
|
+
router.get(f"{HOST}/api/2.0/sql/warehouses/{WAREHOUSE_ID}").mock(
|
|
58
|
+
side_effect=_slow_response(_WAREHOUSE_CHECK_COST_S, {"state": "RUNNING"})
|
|
59
|
+
)
|
|
60
|
+
router.post(f"{HOST}/api/2.0/sql/statements").mock(
|
|
61
|
+
return_value=httpx.Response(
|
|
62
|
+
200,
|
|
63
|
+
json={
|
|
64
|
+
"statement_id": statement_id,
|
|
65
|
+
"status": {"state": "SUCCEEDED"},
|
|
66
|
+
"manifest": {"chunks": [{"chunk_index": 0, "row_count": 1}]},
|
|
67
|
+
},
|
|
68
|
+
)
|
|
69
|
+
)
|
|
70
|
+
router.get(url__regex=rf"{HOST}/api/2\.0/sql/statements/{statement_id}/result/chunks/\d+").mock(
|
|
71
|
+
return_value=httpx.Response(200, json={"external_links": [{"external_link": f"{HOST}/_data/chunk-0"}]})
|
|
72
|
+
)
|
|
73
|
+
router.get(url__regex=rf"{HOST}/_data/chunk-\d+").mock(
|
|
74
|
+
return_value=httpx.Response(200, content=build_chunk_bytes(0, 1))
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def _slow_response(delay_s: float, body: dict) -> object:
|
|
79
|
+
async def _handler(request: httpx.Request) -> httpx.Response:
|
|
80
|
+
await asyncio.sleep(delay_s)
|
|
81
|
+
return httpx.Response(200, json=body)
|
|
82
|
+
|
|
83
|
+
return _handler
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
async def _timed_query(client: DatabricksClient, sql: str) -> float:
|
|
87
|
+
"""Charges _HANDSHAKE_COST_S the first time this client's shared
|
|
88
|
+
httpx.AsyncClient gets its first real use, simulating a fresh
|
|
89
|
+
connection's TCP+TLS setup -- a real socket would pay this on its own,
|
|
90
|
+
respx's in-memory transport doesn't, so it's added explicitly here."""
|
|
91
|
+
is_first_use = client._http is None
|
|
92
|
+
start = time.perf_counter()
|
|
93
|
+
if is_first_use:
|
|
94
|
+
await asyncio.sleep(_HANDSHAKE_COST_S)
|
|
95
|
+
await client.execute_json_statement(sql)
|
|
96
|
+
return time.perf_counter() - start
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
async def _run_cold(n: int, sql: str) -> list[float]:
|
|
100
|
+
times = []
|
|
101
|
+
with respx.mock:
|
|
102
|
+
_install_routes(respx.mock, "stmt-cold")
|
|
103
|
+
for _ in range(n):
|
|
104
|
+
token = "fake-token" # noqa: S105 -- fake token, mocked transport only
|
|
105
|
+
client = DatabricksClient(HOST, WAREHOUSE_ID, token=token)
|
|
106
|
+
times.append(await _timed_query(client, sql))
|
|
107
|
+
await client.aclose()
|
|
108
|
+
return times
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
async def _run_warm(n: int, sql: str) -> list[float]:
|
|
112
|
+
times = []
|
|
113
|
+
with respx.mock:
|
|
114
|
+
_install_routes(respx.mock, "stmt-warm")
|
|
115
|
+
token = "fake-token" # noqa: S105 -- fake token, mocked transport only
|
|
116
|
+
async with DatabricksClient(HOST, WAREHOUSE_ID, token=token) as client:
|
|
117
|
+
for _ in range(n):
|
|
118
|
+
times.append(await _timed_query(client, sql))
|
|
119
|
+
return times
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def _report(label: str, times: list[float]) -> None:
|
|
123
|
+
print(
|
|
124
|
+
f"{label}: n={len(times)} mean={statistics.mean(times) * 1000:.1f}ms "
|
|
125
|
+
f"median={statistics.median(times) * 1000:.1f}ms total={sum(times):.3f}s"
|
|
126
|
+
)
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
async def main() -> None:
|
|
130
|
+
parser = argparse.ArgumentParser()
|
|
131
|
+
parser.add_argument("--n", type=int, default=20)
|
|
132
|
+
args = parser.parse_args()
|
|
133
|
+
|
|
134
|
+
print("*** SIMULATED benchmark -- respx-mocked transport, no real Databricks call. ***")
|
|
135
|
+
print(f"Injected costs: fresh-connection handshake={_HANDSHAKE_COST_S * 1000:.0f}ms, "
|
|
136
|
+
f"warehouse-status round-trip={_WAREHOUSE_CHECK_COST_S * 1000:.0f}ms\n")
|
|
137
|
+
|
|
138
|
+
cold = await _run_cold(args.n, "SELECT 1")
|
|
139
|
+
_report("cold (old behavior: fresh client + warehouse-check per query)", cold)
|
|
140
|
+
|
|
141
|
+
warm = await _run_warm(args.n, "SELECT 1")
|
|
142
|
+
_report("warm (new behavior: one reused client) ", warm)
|
|
143
|
+
|
|
144
|
+
saved_per_query_ms = (statistics.mean(cold) - statistics.mean(warm)) * 1000
|
|
145
|
+
speedup = statistics.mean(cold) / statistics.mean(warm) if statistics.mean(warm) > 0 else float("inf")
|
|
146
|
+
print(f"\nmean saved per query (after the first): ~{saved_per_query_ms:.1f}ms")
|
|
147
|
+
print(f"mean speedup: {speedup:.2f}x")
|
|
148
|
+
print(
|
|
149
|
+
"\nNote: this isolates exactly the two mechanisms changed (connection reuse, "
|
|
150
|
+
"warehouse-check caching) with realistic but assumed latency values -- it does "
|
|
151
|
+
"not measure real Databricks/network round-trip times, statement execution time, "
|
|
152
|
+
"or chunk-fetch time, none of which changed."
|
|
153
|
+
)
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
if __name__ == "__main__":
|
|
157
|
+
asyncio.run(main())
|
|
@@ -82,6 +82,7 @@ class DatabricksClient:
|
|
|
82
82
|
wait_timeout: str = "30s",
|
|
83
83
|
chunk_fetch_concurrency: int = 6,
|
|
84
84
|
warehouse_start_timeout: float = 300.0,
|
|
85
|
+
warehouse_confirmed_running_ttl_s: float = 30.0,
|
|
85
86
|
) -> None:
|
|
86
87
|
if not token and not token_provider:
|
|
87
88
|
raise ValueError("DatabricksClient needs either `token` or `token_provider`")
|
|
@@ -101,6 +102,40 @@ class DatabricksClient:
|
|
|
101
102
|
# slow consumer.
|
|
102
103
|
self.chunk_fetch_concurrency = chunk_fetch_concurrency
|
|
103
104
|
self.warehouse_start_timeout = warehouse_start_timeout
|
|
105
|
+
# See _ensure_warehouse_running: a warm warehouse doesn't need its
|
|
106
|
+
# RUNNING state re-verified on every single statement -- one round
|
|
107
|
+
# trip saved per query once confirmed within this window.
|
|
108
|
+
self._warehouse_confirmed_running_ttl_s = warehouse_confirmed_running_ttl_s
|
|
109
|
+
self._warehouse_confirmed_running_at: float | None = None
|
|
110
|
+
# One shared connection pool for this client's whole lifetime instead
|
|
111
|
+
# of a fresh httpx.AsyncClient() (and its own TCP+TLS handshake) per
|
|
112
|
+
# call -- every method below hits the same Databricks host repeatedly,
|
|
113
|
+
# so keep-alive/pooling actually pays off across calls. Presigned
|
|
114
|
+
# external-link downloads (a different host: blob storage) still each
|
|
115
|
+
# get their own connection from this same pool as needed; httpx pools
|
|
116
|
+
# per-host internally, so sharing one client across hosts is safe.
|
|
117
|
+
self._http: httpx.AsyncClient | None = None
|
|
118
|
+
self._http_lock = asyncio.Lock()
|
|
119
|
+
|
|
120
|
+
async def _get_http_client(self) -> httpx.AsyncClient:
|
|
121
|
+
if self._http is None:
|
|
122
|
+
async with self._http_lock:
|
|
123
|
+
if self._http is None:
|
|
124
|
+
self._http = httpx.AsyncClient()
|
|
125
|
+
return self._http
|
|
126
|
+
|
|
127
|
+
async def aclose(self) -> None:
|
|
128
|
+
"""Closes the shared connection pool. Safe to call even if no request
|
|
129
|
+
was ever made (no-op) or more than once."""
|
|
130
|
+
if self._http is not None:
|
|
131
|
+
await self._http.aclose()
|
|
132
|
+
self._http = None
|
|
133
|
+
|
|
134
|
+
async def __aenter__(self) -> DatabricksClient:
|
|
135
|
+
return self
|
|
136
|
+
|
|
137
|
+
async def __aexit__(self, *exc_info: object) -> None:
|
|
138
|
+
await self.aclose()
|
|
104
139
|
|
|
105
140
|
async def _bearer_token(self) -> str:
|
|
106
141
|
if self._token is not None:
|
|
@@ -135,12 +170,23 @@ class DatabricksClient:
|
|
|
135
170
|
implicit auto-start. A cold warehouse's catalog credential cache needs
|
|
136
171
|
a moment to catch up right after startup -- submitting straight into
|
|
137
172
|
that window is a common source of transient, identity-scoped 403s.
|
|
138
|
-
|
|
139
|
-
|
|
173
|
+
|
|
174
|
+
Skips the check entirely if RUNNING was already confirmed within
|
|
175
|
+
`_warehouse_confirmed_running_ttl_s` -- a warm, always-on warehouse
|
|
176
|
+
doesn't need re-verifying on every single statement; that GET is a
|
|
177
|
+
full round trip that buys nothing once already known-good."""
|
|
178
|
+
now = time.monotonic()
|
|
179
|
+
if (
|
|
180
|
+
self._warehouse_confirmed_running_at is not None
|
|
181
|
+
and now - self._warehouse_confirmed_running_at < self._warehouse_confirmed_running_ttl_s
|
|
182
|
+
):
|
|
183
|
+
return
|
|
184
|
+
|
|
140
185
|
url = f"{self._host}/api/2.0/sql/warehouses/{self.warehouse_id}"
|
|
141
186
|
resp = await self._authed_request(client, "GET", url)
|
|
142
187
|
state = resp.json().get("state")
|
|
143
188
|
if state == "RUNNING":
|
|
189
|
+
self._warehouse_confirmed_running_at = time.monotonic()
|
|
144
190
|
return
|
|
145
191
|
|
|
146
192
|
if state == "STOPPED":
|
|
@@ -152,6 +198,7 @@ class DatabricksClient:
|
|
|
152
198
|
resp = await self._authed_request(client, "GET", url)
|
|
153
199
|
state = resp.json().get("state")
|
|
154
200
|
if state == "RUNNING":
|
|
201
|
+
self._warehouse_confirmed_running_at = time.monotonic()
|
|
155
202
|
return
|
|
156
203
|
# Falls through and lets the statement submission itself surface
|
|
157
204
|
# whatever's actually wrong -- proceeding anyway rather than raising
|
|
@@ -187,21 +234,21 @@ class DatabricksClient:
|
|
|
187
234
|
if parameters:
|
|
188
235
|
body["parameters"] = parameters
|
|
189
236
|
|
|
190
|
-
|
|
191
|
-
|
|
192
|
-
|
|
193
|
-
|
|
237
|
+
client = await self._get_http_client()
|
|
238
|
+
await self._ensure_warehouse_running(client)
|
|
239
|
+
resp = await self._authed_request(client, "POST", f"{self._host}/api/2.0/sql/statements", json=body)
|
|
240
|
+
data = resp.json()
|
|
194
241
|
|
|
242
|
+
status = data.get("status", {})
|
|
243
|
+
while status.get("state") not in _TERMINAL_STATES:
|
|
244
|
+
statement_id = data["statement_id"]
|
|
245
|
+
await asyncio.sleep(_POLL_INTERVAL_S)
|
|
246
|
+
resp = await self._authed_request(client, "GET", f"{self._host}/api/2.0/sql/statements/{statement_id}")
|
|
247
|
+
data = resp.json()
|
|
195
248
|
status = data.get("status", {})
|
|
196
|
-
while status.get("state") not in _TERMINAL_STATES:
|
|
197
|
-
statement_id = data["statement_id"]
|
|
198
|
-
await asyncio.sleep(_POLL_INTERVAL_S)
|
|
199
|
-
resp = await self._authed_request(client, "GET", f"{self._host}/api/2.0/sql/statements/{statement_id}")
|
|
200
|
-
data = resp.json()
|
|
201
|
-
status = data.get("status", {})
|
|
202
249
|
|
|
203
|
-
|
|
204
|
-
|
|
250
|
+
_raise_for_failed(status)
|
|
251
|
+
return data["statement_id"], data.get("manifest") or {}
|
|
205
252
|
|
|
206
253
|
async def execute_arrow_statement(
|
|
207
254
|
self,
|
|
@@ -249,26 +296,26 @@ class DatabricksClient:
|
|
|
249
296
|
overwriting anything already there. `volume_path` is caller-supplied
|
|
250
297
|
in full (e.g. `/Volumes/my_catalog/my_schema/my_volume/some/file.parquet`)
|
|
251
298
|
-- this package has no knowledge of any specific catalog/schema/volume."""
|
|
252
|
-
|
|
253
|
-
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
299
|
+
client = await self._get_http_client()
|
|
300
|
+
await self._authed_request(
|
|
301
|
+
client,
|
|
302
|
+
"PUT",
|
|
303
|
+
f"{self._host}/api/2.0/fs/files{volume_path}",
|
|
304
|
+
params={"overwrite": "true"},
|
|
305
|
+
content_type="application/octet-stream",
|
|
306
|
+
content=data,
|
|
307
|
+
)
|
|
261
308
|
|
|
262
309
|
async def delete_volume_file(self, volume_path: str) -> None:
|
|
263
310
|
"""Deletes a file at `volume_path` (see upload_volume_file). A 404 is
|
|
264
311
|
treated as success -- the file is already gone, which is fine for
|
|
265
312
|
idempotent staging cleanup."""
|
|
266
|
-
|
|
267
|
-
|
|
268
|
-
|
|
269
|
-
|
|
270
|
-
|
|
271
|
-
|
|
313
|
+
client = await self._get_http_client()
|
|
314
|
+
try:
|
|
315
|
+
await self._authed_request(client, "DELETE", f"{self._host}/api/2.0/fs/files{volume_path}")
|
|
316
|
+
except httpx.HTTPStatusError as exc:
|
|
317
|
+
if exc.response.status_code != 404:
|
|
318
|
+
raise
|
|
272
319
|
|
|
273
320
|
async def stream_chunks_by_index(
|
|
274
321
|
self, statement_id: str, chunk_metas: list[dict[str, Any]]
|
|
@@ -281,11 +328,11 @@ class DatabricksClient:
|
|
|
281
328
|
manifest didn't carry one) and its own chunk_index, so a caller that
|
|
282
329
|
cares about the original row order (e.g. a query with ORDER BY) can
|
|
283
330
|
restore it even though chunks can complete out of order."""
|
|
284
|
-
|
|
285
|
-
|
|
286
|
-
|
|
287
|
-
|
|
288
|
-
|
|
331
|
+
client = await self._get_http_client()
|
|
332
|
+
async for blob, row_count, chunk_index in self._fetch_chunks_with_backpressure(
|
|
333
|
+
client, statement_id, chunk_metas
|
|
334
|
+
):
|
|
335
|
+
yield blob, row_count, chunk_index
|
|
289
336
|
|
|
290
337
|
async def _fetch_link_bytes(self, client: httpx.AsyncClient, url: str) -> bytes:
|
|
291
338
|
async def _do() -> bytes:
|
|
@@ -47,12 +47,14 @@ def mock_warehouse():
|
|
|
47
47
|
Callers configure it via mock_warehouse(...) inside a `with respx.mock:`
|
|
48
48
|
block (or use the `respx_router` param)."""
|
|
49
49
|
|
|
50
|
-
def _install(
|
|
50
|
+
def _install(
|
|
51
|
+
router: respx.Router, n_chunks: int, rows_per_chunk: int, *, reverse_arrival: bool = False
|
|
52
|
+
) -> respx.Route:
|
|
51
53
|
statement_id = "stmt-abc"
|
|
52
54
|
chunks = [{"chunk_index": i, "row_count": rows_per_chunk} for i in range(n_chunks)]
|
|
53
55
|
chunk_bytes = {i: build_chunk_bytes(i * rows_per_chunk, (i + 1) * rows_per_chunk) for i in range(n_chunks)}
|
|
54
56
|
|
|
55
|
-
router.get(f"{HOST}/api/2.0/sql/warehouses/{WAREHOUSE_ID}").mock(
|
|
57
|
+
warehouse_route = router.get(f"{HOST}/api/2.0/sql/warehouses/{WAREHOUSE_ID}").mock(
|
|
56
58
|
return_value=httpx.Response(200, json={"state": "RUNNING"})
|
|
57
59
|
)
|
|
58
60
|
router.post(f"{HOST}/api/2.0/sql/statements").mock(
|
|
@@ -89,6 +91,7 @@ def mock_warehouse():
|
|
|
89
91
|
return httpx.Response(200, content=chunk_bytes[idx])
|
|
90
92
|
|
|
91
93
|
router.get(url__regex=rf"{HOST}/_data/chunk-\d+").mock(side_effect=serve_chunk_bytes)
|
|
94
|
+
return warehouse_route
|
|
92
95
|
|
|
93
96
|
return _install
|
|
94
97
|
|
|
@@ -163,3 +163,84 @@ def test_requires_token_or_provider(warehouse_host_id):
|
|
|
163
163
|
host, warehouse_id = warehouse_host_id
|
|
164
164
|
with pytest.raises(ValueError, match="token"):
|
|
165
165
|
DatabricksClient(host, warehouse_id)
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
@pytest.mark.asyncio
|
|
169
|
+
@respx.mock
|
|
170
|
+
async def test_shares_one_http_client_across_calls(mock_warehouse, warehouse_host_id):
|
|
171
|
+
"""A DatabricksClient should reuse one httpx.AsyncClient (and its
|
|
172
|
+
connection pool) across statements and chunk fetches, not open a fresh
|
|
173
|
+
one per call -- see client.py's _get_http_client."""
|
|
174
|
+
host, warehouse_id = warehouse_host_id
|
|
175
|
+
mock_warehouse(respx.mock, n_chunks=2, rows_per_chunk=2)
|
|
176
|
+
client = DatabricksClient(host, warehouse_id, token="test-token")
|
|
177
|
+
|
|
178
|
+
assert client._http is None
|
|
179
|
+
await client.execute_json_statement("SELECT 1")
|
|
180
|
+
first = client._http
|
|
181
|
+
assert first is not None
|
|
182
|
+
await client.execute_json_statement("SELECT 2")
|
|
183
|
+
assert client._http is first # same instance reused, not recreated
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
@pytest.mark.asyncio
|
|
187
|
+
@respx.mock
|
|
188
|
+
async def test_aclose_closes_and_allows_reopening(mock_warehouse, warehouse_host_id):
|
|
189
|
+
host, warehouse_id = warehouse_host_id
|
|
190
|
+
mock_warehouse(respx.mock, n_chunks=1, rows_per_chunk=1)
|
|
191
|
+
client = DatabricksClient(host, warehouse_id, token="test-token")
|
|
192
|
+
|
|
193
|
+
await client.execute_json_statement("SELECT 1")
|
|
194
|
+
first = client._http
|
|
195
|
+
assert first is not None
|
|
196
|
+
await client.aclose()
|
|
197
|
+
assert client._http is None
|
|
198
|
+
assert first.is_closed
|
|
199
|
+
|
|
200
|
+
await client.execute_json_statement("SELECT 1") # still usable after close
|
|
201
|
+
assert client._http is not None
|
|
202
|
+
assert client._http is not first
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
@pytest.mark.asyncio
|
|
206
|
+
@respx.mock
|
|
207
|
+
async def test_async_context_manager_closes_on_exit(mock_warehouse, warehouse_host_id):
|
|
208
|
+
host, warehouse_id = warehouse_host_id
|
|
209
|
+
mock_warehouse(respx.mock, n_chunks=1, rows_per_chunk=1)
|
|
210
|
+
|
|
211
|
+
async with DatabricksClient(host, warehouse_id, token="test-token") as client:
|
|
212
|
+
await client.execute_json_statement("SELECT 1")
|
|
213
|
+
http_client = client._http
|
|
214
|
+
|
|
215
|
+
assert http_client is not None
|
|
216
|
+
assert http_client.is_closed
|
|
217
|
+
assert client._http is None
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
@pytest.mark.asyncio
|
|
221
|
+
@respx.mock
|
|
222
|
+
async def test_warehouse_running_check_is_cached_across_statements(mock_warehouse, warehouse_host_id):
|
|
223
|
+
"""Once RUNNING is confirmed, a second statement within the TTL window
|
|
224
|
+
shouldn't re-GET the warehouse status -- that round trip buys nothing on
|
|
225
|
+
an already-known-warm warehouse (see _ensure_warehouse_running)."""
|
|
226
|
+
host, warehouse_id = warehouse_host_id
|
|
227
|
+
warehouse_route = mock_warehouse(respx.mock, n_chunks=1, rows_per_chunk=1)
|
|
228
|
+
client = DatabricksClient(host, warehouse_id, token="test-token", warehouse_confirmed_running_ttl_s=60.0)
|
|
229
|
+
|
|
230
|
+
await client.execute_json_statement("SELECT 1")
|
|
231
|
+
assert warehouse_route.call_count == 1
|
|
232
|
+
await client.execute_json_statement("SELECT 2")
|
|
233
|
+
assert warehouse_route.call_count == 1 # still 1 -- cached, no second GET
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
@pytest.mark.asyncio
|
|
237
|
+
@respx.mock
|
|
238
|
+
async def test_warehouse_running_check_re_verifies_after_ttl_expires(mock_warehouse, warehouse_host_id):
|
|
239
|
+
host, warehouse_id = warehouse_host_id
|
|
240
|
+
warehouse_route = mock_warehouse(respx.mock, n_chunks=1, rows_per_chunk=1)
|
|
241
|
+
client = DatabricksClient(host, warehouse_id, token="test-token", warehouse_confirmed_running_ttl_s=0.0)
|
|
242
|
+
|
|
243
|
+
await client.execute_json_statement("SELECT 1")
|
|
244
|
+
assert warehouse_route.call_count == 1
|
|
245
|
+
await client.execute_json_statement("SELECT 2")
|
|
246
|
+
assert warehouse_route.call_count == 2 # TTL is 0 -- re-verified every time
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|