arrowbricks 0.1.2__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.
Files changed (27) hide show
  1. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/AGENTS.md +4 -3
  2. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/PKG-INFO +8 -4
  3. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/README.md +7 -3
  4. arrowbricks-0.3.0/examples/fastapi_sse_pivot.py +78 -0
  5. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/pyproject.toml +1 -1
  6. arrowbricks-0.3.0/scripts/benchmark_client_reuse.py +126 -0
  7. arrowbricks-0.3.0/scripts/benchmark_simulated.py +157 -0
  8. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/src/arrowbricks/_streaming.py +26 -8
  9. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/src/arrowbricks/client.py +81 -34
  10. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/src/arrowbricks/cursor.py +61 -11
  11. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/tests/conftest.py +5 -2
  12. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/tests/test_client.py +81 -0
  13. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/tests/test_cursor.py +57 -1
  14. arrowbricks-0.3.0/tests/test_result_set.py +74 -0
  15. arrowbricks-0.3.0/tests/test_streaming.py +82 -0
  16. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/uv.lock +1 -1
  17. arrowbricks-0.1.2/tests/test_streaming.py +0 -39
  18. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/.github/workflows/ci.yml +0 -0
  19. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/.github/workflows/release.yml +0 -0
  20. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/.gitignore +0 -0
  21. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/LICENSE +0 -0
  22. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/examples/azure_auth.py +0 -0
  23. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/examples/basic.py +0 -0
  24. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/examples/cursor_paging.py +0 -0
  25. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/examples/fastapi_sse.py +0 -0
  26. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/prek.toml +0 -0
  27. {arrowbricks-0.1.2 → arrowbricks-0.3.0}/src/arrowbricks/__init__.py +0 -0
@@ -6,9 +6,9 @@ Runs SQL against a Databricks SQL warehouse via the Statement Execution API and
6
6
 
7
7
  - `src/arrowbricks/client.py` -- pure REST client (auth, statement submission/polling, backpressure-bounded concurrent chunk download). No Arrow dependency at all, intentionally: someone who only wants `execute_json_statement` or `upload_volume_file`/`delete_volume_file` shouldn't need arro3 pulled in conceptually either, even though it's a hard dependency of the package as a whole. Retries are a small hand-rolled `_retry_call` loop, not a dependency (see "Design invariants" below).
8
8
  - `src/arrowbricks/_streaming.py` -- Arrow (de)serialization via arro3, always: `ReplayableArrowChunk`, `write_ipc_stream` (always uncompressed, see below), heartbeat helpers, chunk fetching, and `stream_query_json` (arro3's `write_ndjson`). One Arrow engine, no pluggable backend -- unlike duckbricks, there's nothing here to make pluggable; arro3 *is* the whole point.
9
- - `src/arrowbricks/cursor.py` -- `Connection`/`Cursor`, the DB-API-ish surface (`execute`/`execute_streamed`, `fetchone`/`fetchmany`/`fetchall`, `fetchall_arrow`/`fetchmany_arrow`). `_ResultSet` buffers at the Arrow-Table level (not materialized Python rows) so the Arrow-native fetch methods stay zero-copy; row-based fetches materialize lazily off that buffer.
9
+ - `src/arrowbricks/cursor.py` -- `Connection`/`Cursor`, the DB-API-ish surface (`execute`/`execute_streamed`, `fetchone`/`fetchmany`/`fetchall`, `fetchall_arrow`/`fetchmany_arrow`, `fetchall_streamed`/`fetchall_arrow_streamed`). `_ResultSet` buffers at the Arrow-Table level (not materialized Python rows) so the Arrow-native fetch methods stay zero-copy; row-based fetches materialize lazily off that buffer.
10
10
  - `tests/` -- respx mocks the Databricks REST endpoints (warehouse status, statement submit, chunk-link resolution, external-link byte download); no real warehouse or credentials needed to run the suite.
11
- - `examples/` -- `basic.py` (static token), `cursor_paging.py` (fetchmany/fetchmany_arrow over a large result), `fastapi_sse.py` (streaming NDJSON as SSE), `azure_auth.py` (Azure AD `token_provider` via `azure-identity`, kept out of core deps on purpose -- don't widen `ty check`'s scope to include it).
11
+ - `examples/` -- `basic.py` (static token), `cursor_paging.py` (fetchmany/fetchmany_arrow over a large result), `fastapi_sse.py` (streaming NDJSON as SSE), `fastapi_sse_pivot.py` (buffered fetchall_streamed with one combined heartbeat/timeout budget across the wait+download phases), `azure_auth.py` (Azure AD `token_provider` via `azure-identity`, kept out of core deps on purpose -- don't widen `ty check`'s scope to include it).
12
12
 
13
13
  ## Commands
14
14
 
@@ -25,8 +25,9 @@ One-time setup per clone: `prek install` (needs `uv tool install prek` first if
25
25
 
26
26
  - **No cloud-SDK dependency.** Auth is `token: str` or `token_provider: Callable[[], str | Awaitable[str]]`. Do not add `azure-identity`/`boto3`/etc. as a real dependency -- that belongs in the caller's app.
27
27
  - **No hardcoded catalog/schema.** `catalog`/`schema` default to `None` everywhere. This package has zero knowledge of any specific Databricks workspace's naming.
28
- - **Chunk order is not fetch order.** `DatabricksClient` fetches chunks concurrently (bounded, with backpressure) and they can complete out of order. `_ResultSet`/`stream_query_json` both hold a `pending: dict[int, chunk]` reorder buffer keyed by `chunk_index`, releasing in order as the next expected index shows up. If you touch either, keep a test proving order survives out-of-order arrival (see `test_fetchall_preserves_order_despite_out_of_order_chunks`, `test_stream_query_json_preserves_order_despite_out_of_order_chunks`).
28
+ - **Chunk order is not fetch order, and a chunk_index is not guaranteed unique or contiguous.** `DatabricksClient` fetches chunks concurrently (bounded, with backpressure) and they can complete out of order. `_ResultSet`/`stream_query_json` both hold a `pending: dict[int, list[chunk]]` reorder buffer keyed by `chunk_index` -- a **list**, not a single chunk, because `_fetch_chunk_index` can yield more than one blob for the same index (multiple `external_links` per chunk, "usually exactly one" but not guaranteed). Releases in order as the next expected index empties out; once the source is exhausted, a genuine gap (an index that never showed up at all) drains the lowest remaining index instead of stranding everything buffered past it. If you touch either, keep a test proving order survives out-of-order arrival AND that duplicate/missing indices never lose rows (see `test_fetchall_preserves_order_despite_out_of_order_chunks`, `test_stream_query_json_preserves_order_despite_out_of_order_chunks`, `tests/test_result_set.py`).
29
29
  - **Chunks are fetched lazily, not all upfront.** `_ResultSet` only pulls the next chunk from `_chunk_aiter` when the caller's `fetchone`/`fetchmany`/`fetchall` actually needs more rows than are already buffered. Don't "simplify" this into draining the whole chunk iterator inside `execute()` -- that defeats the point of a paginated cursor.
30
+ - **`execute_streamed`'s heartbeat/timeout only covers the wait for the statement to become ready -- not downloading any chunk.** That was a real bug in this package's first release: a caller doing `execute_streamed()` then `fetchall()` got zero timeout enforcement and zero heartbeats during a slow multi-chunk download, exactly the case heartbeats exist for. `fetchall_streamed`/`fetchall_arrow_streamed` cover that second phase -- a caller wanting one combined budget across both phases must track its own deadline and pass the *remaining* time into the second call (see `examples/fastapi_sse_pivot.py`), since two independently-clocked `total_timeout_s`s would let a pathological case run up to 2x the intended ceiling.
30
31
  - **No silent row caps.** There's no `ABSOLUTE_ROW_LIMIT`-style ceiling baked in. If a caller wants one, that's `row_limit`, which they pass explicitly.
31
32
  - **No retry dependency.** `client.py`'s `_retry_call` is a ~10-line hand-rolled exponential-backoff loop, replacing tenacity on purpose -- it's the only retry pattern in the whole client, so a dependency for it wasn't worth it.
32
33
  - **One Arrow engine, no pluggable backend.** Unlike duckbricks (which supports nanoarrow *or* arro3 *or* bring-your-own), arrowbricks is arro3-only by design -- that's the entire "single responsibility" pitch. Don't add a backend-abstraction layer back in; if a caller needs a different Arrow engine, that's duckbricks' `set_arrow_backend()`, not this package.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: arrowbricks
3
- Version: 0.1.2
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
@@ -72,8 +72,11 @@ See [`examples/basic.py`](examples/basic.py) for a runnable version,
72
72
  [`examples/cursor_paging.py`](examples/cursor_paging.py) for paging a large
73
73
  result with `fetchmany`/`fetchmany_arrow` without buffering it all upfront,
74
74
  [`examples/fastapi_sse.py`](examples/fastapi_sse.py) for streaming a query to
75
- a client as Server-Sent Events, or [`examples/azure_auth.py`](examples/azure_auth.py)
76
- for a caching `token_provider` built on Azure AD (`DefaultAzureCredential`).
75
+ a client as Server-Sent Events, [`examples/fastapi_sse_pivot.py`](examples/fastapi_sse_pivot.py)
76
+ for the same over a buffered `Cursor.fetchall_streamed` result with one
77
+ combined heartbeat/timeout budget across both the wait and the download, or
78
+ [`examples/azure_auth.py`](examples/azure_auth.py) for a caching
79
+ `token_provider` built on Azure AD (`DefaultAzureCredential`).
77
80
 
78
81
  ## Why not `databricks-sql-connector`?
79
82
 
@@ -102,9 +105,10 @@ conn = connect(host=..., warehouse_id=..., token_provider=my_token_provider)
102
105
  - `Connection.cursor() -> Cursor`
103
106
  - `Connection.client -> DatabricksClient` -- the same client `cursor()` uses, for lower-level access (e.g. `stream_query_json`, `execute_json_statement`, `upload_volume_file`).
104
107
  - `Cursor.execute(sql, parameters=None, *, row_limit=None, offset=None, catalog=None, schema=None, total_timeout_s=None) -> Cursor` -- submits and waits for the statement, like a real DB-API cursor. `parameters`, if given, is Databricks' own named-parameter format -- `[{"name": ..., "value": ..., "type": ...}]` bound against `:name` markers in `sql`.
105
- - `Cursor.execute_streamed(...)` -- same args, but an async generator yielding `HEARTBEAT` while waiting on a slow cold start, then the ready `Cursor` -- for bridging e.g. an SSE connection.
108
+ - `Cursor.execute_streamed(...)` -- same args, but an async generator yielding `HEARTBEAT` while waiting on a slow cold start, then the ready `Cursor` -- for bridging e.g. an SSE connection. Its timeout/heartbeats stop the moment the statement is ready, *before* any chunk has been downloaded -- see `fetchall_streamed` below for the download phase itself.
106
109
  - `Cursor.fetchone() -> tuple | None`, `Cursor.fetchmany(size) -> list[tuple]`, `Cursor.fetchall() -> list[tuple]`
107
110
  - `Cursor.fetchmany_arrow(size) -> arro3.core.Table`, `Cursor.fetchall_arrow() -> arro3.core.Table`
111
+ - `Cursor.fetchall_streamed(*, total_timeout_s=None)` / `Cursor.fetchall_arrow_streamed(*, total_timeout_s=None)` -- like `fetchall()`/`fetchall_arrow()`, but yield `HEARTBEAT` while pulling chunks instead of blocking silently, then the final rows/Table -- for a caller downloading a large result over SSE who needs heartbeats (and a timeout) through the *download*, not just the initial wait. Compose with `execute_streamed` and a shared deadline if you want one combined budget across both phases (see `examples/fastapi_sse_pivot.py`).
108
112
  - `Cursor` is an async iterator, yielding one row (tuple) at a time.
109
113
  - `Cursor.description` -- DB-API-style `[(name, type_name, None, None, None, None, None), ...]` after `execute()`.
110
114
  - `stream_query_json(client, sql, **kwargs)` -- yields `HEARTBEAT`, then each row as a JSON string, as soon as its chunk arrives. Timestamps come out as full ISO-8601, every column key is always present (`"col":null` for a null value, never an omitted key).
@@ -59,8 +59,11 @@ See [`examples/basic.py`](examples/basic.py) for a runnable version,
59
59
  [`examples/cursor_paging.py`](examples/cursor_paging.py) for paging a large
60
60
  result with `fetchmany`/`fetchmany_arrow` without buffering it all upfront,
61
61
  [`examples/fastapi_sse.py`](examples/fastapi_sse.py) for streaming a query to
62
- a client as Server-Sent Events, or [`examples/azure_auth.py`](examples/azure_auth.py)
63
- for a caching `token_provider` built on Azure AD (`DefaultAzureCredential`).
62
+ a client as Server-Sent Events, [`examples/fastapi_sse_pivot.py`](examples/fastapi_sse_pivot.py)
63
+ for the same over a buffered `Cursor.fetchall_streamed` result with one
64
+ combined heartbeat/timeout budget across both the wait and the download, or
65
+ [`examples/azure_auth.py`](examples/azure_auth.py) for a caching
66
+ `token_provider` built on Azure AD (`DefaultAzureCredential`).
64
67
 
65
68
  ## Why not `databricks-sql-connector`?
66
69
 
@@ -89,9 +92,10 @@ conn = connect(host=..., warehouse_id=..., token_provider=my_token_provider)
89
92
  - `Connection.cursor() -> Cursor`
90
93
  - `Connection.client -> DatabricksClient` -- the same client `cursor()` uses, for lower-level access (e.g. `stream_query_json`, `execute_json_statement`, `upload_volume_file`).
91
94
  - `Cursor.execute(sql, parameters=None, *, row_limit=None, offset=None, catalog=None, schema=None, total_timeout_s=None) -> Cursor` -- submits and waits for the statement, like a real DB-API cursor. `parameters`, if given, is Databricks' own named-parameter format -- `[{"name": ..., "value": ..., "type": ...}]` bound against `:name` markers in `sql`.
92
- - `Cursor.execute_streamed(...)` -- same args, but an async generator yielding `HEARTBEAT` while waiting on a slow cold start, then the ready `Cursor` -- for bridging e.g. an SSE connection.
95
+ - `Cursor.execute_streamed(...)` -- same args, but an async generator yielding `HEARTBEAT` while waiting on a slow cold start, then the ready `Cursor` -- for bridging e.g. an SSE connection. Its timeout/heartbeats stop the moment the statement is ready, *before* any chunk has been downloaded -- see `fetchall_streamed` below for the download phase itself.
93
96
  - `Cursor.fetchone() -> tuple | None`, `Cursor.fetchmany(size) -> list[tuple]`, `Cursor.fetchall() -> list[tuple]`
94
97
  - `Cursor.fetchmany_arrow(size) -> arro3.core.Table`, `Cursor.fetchall_arrow() -> arro3.core.Table`
98
+ - `Cursor.fetchall_streamed(*, total_timeout_s=None)` / `Cursor.fetchall_arrow_streamed(*, total_timeout_s=None)` -- like `fetchall()`/`fetchall_arrow()`, but yield `HEARTBEAT` while pulling chunks instead of blocking silently, then the final rows/Table -- for a caller downloading a large result over SSE who needs heartbeats (and a timeout) through the *download*, not just the initial wait. Compose with `execute_streamed` and a shared deadline if you want one combined budget across both phases (see `examples/fastapi_sse_pivot.py`).
95
99
  - `Cursor` is an async iterator, yielding one row (tuple) at a time.
96
100
  - `Cursor.description` -- DB-API-style `[(name, type_name, None, None, None, None, None), ...]` after `execute()`.
97
101
  - `stream_query_json(client, sql, **kwargs)` -- yields `HEARTBEAT`, then each row as a JSON string, as soon as its chunk arrives. Timestamps come out as full ISO-8601, every column key is always present (`"col":null` for a null value, never an omitted key).
@@ -0,0 +1,78 @@
1
+ """Stream a large result to a client over SSE with ONE combined heartbeat/
2
+ timeout budget spanning both phases: waiting for the statement to complete
3
+ (execute_streamed) AND downloading every chunk afterwards (fetchall_streamed)
4
+ -- the two-phase composition `fastapi_sse.py`'s simpler stream_query_json
5
+ example doesn't need, since that one streams rows as chunks arrive rather
6
+ than buffering a full result first.
7
+
8
+ Splitting the timeout into two separately-clocked halves (each given the
9
+ same total_timeout_s) would let a pathological case run up to 2x the
10
+ intended ceiling -- tracking one deadline and passing the *remaining* budget
11
+ into the second phase keeps it as one real ceiling across both.
12
+
13
+ Requires `pip install arrowbricks fastapi uvicorn`.
14
+
15
+ DATABRICKS_HOST=adb-1234567890.1.azuredatabricks.net \\
16
+ DATABRICKS_WAREHOUSE_ID=abcd1234efgh5678 \\
17
+ DATABRICKS_TOKEN=dapi... \\
18
+ uvicorn examples.fastapi_sse_pivot:app --reload
19
+
20
+ Then, in another terminal:
21
+
22
+ curl -N "http://localhost:8000/pivot?sql=SELECT+*+FROM+range(1000000)"
23
+
24
+ Same caveat as fastapi_sse.py: `sql` comes straight from the request for
25
+ brevity -- validate/allowlist it in a real deployment.
26
+ """
27
+
28
+ import asyncio
29
+ import os
30
+ from collections.abc import AsyncIterator
31
+
32
+ from fastapi import FastAPI
33
+ from fastapi.responses import StreamingResponse
34
+
35
+ from arrowbricks import HEARTBEAT, connect
36
+ from arrowbricks.cursor import Cursor
37
+
38
+ app = FastAPI()
39
+
40
+ conn = connect(
41
+ host=os.environ["DATABRICKS_HOST"],
42
+ warehouse_id=os.environ["DATABRICKS_WAREHOUSE_ID"],
43
+ token=os.environ["DATABRICKS_TOKEN"],
44
+ )
45
+
46
+ _TOTAL_TIMEOUT_S = 300.0
47
+
48
+
49
+ async def _fetch_all_rows_with_heartbeat(cursor: Cursor, deadline: float) -> AsyncIterator[object]:
50
+ """Yields HEARTBEAT while downloading, then the final list[Row] -- bounded
51
+ by whatever's left of the shared deadline, not a fresh full timeout."""
52
+ loop = asyncio.get_running_loop()
53
+ remaining = max(deadline - loop.time(), 0.0)
54
+ async for item in cursor.fetchall_streamed(total_timeout_s=remaining):
55
+ yield item
56
+
57
+
58
+ async def _sse(sql: str) -> AsyncIterator[str]:
59
+ loop = asyncio.get_running_loop()
60
+ deadline = loop.time() + _TOTAL_TIMEOUT_S
61
+ cursor = conn.cursor()
62
+
63
+ async for item in cursor.execute_streamed(sql, total_timeout_s=_TOTAL_TIMEOUT_S):
64
+ if item is HEARTBEAT:
65
+ yield ": keep-alive\n\n"
66
+ continue
67
+ async for fetch_item in _fetch_all_rows_with_heartbeat(cursor, deadline):
68
+ if fetch_item is HEARTBEAT:
69
+ yield ": keep-alive\n\n"
70
+ else:
71
+ for row in fetch_item:
72
+ yield f"data: {row}\n\n"
73
+ yield "event: end\ndata: {}\n\n"
74
+
75
+
76
+ @app.get("/pivot")
77
+ async def pivot(sql: str) -> StreamingResponse:
78
+ return StreamingResponse(_sse(sql), media_type="text/event-stream")
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "arrowbricks"
3
- version = "0.1.2"
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())
@@ -236,7 +236,11 @@ async def stream_query_json(
236
236
  O(whole result). Chunks can arrive out of order, so out-of-order arrivals
237
237
  sit in a small `pending` buffer until the next expected chunk_index shows
238
238
  up -- that buffer stays bounded by concurrency, it never grows to the full
239
- result.
239
+ result. `pending` holds a *list* per index, not a single chunk, since a
240
+ chunk_index can have more than one blob (DatabricksClient._fetch_chunk_index
241
+ gathers over possibly-multiple `external_links` per chunk); once the
242
+ source is exhausted, any index that never showed up at all is skipped
243
+ rather than stranding everything buffered past it forever.
240
244
 
241
245
  Note this yields a whole chunk's rows at once (write_ndjson has no
242
246
  incremental/row-at-a-time mode) -- Databricks' own chunk sizing already
@@ -247,17 +251,31 @@ async def stream_query_json(
247
251
  client, sql, catalog=catalog, schema=schema, parameters=params
248
252
  )
249
253
 
250
- pending: dict[int, ReplayableArrowChunk] = {}
254
+ pending: dict[int, list[ReplayableArrowChunk]] = {}
251
255
  next_idx = 0
252
256
  loop = asyncio.get_running_loop()
257
+
258
+ async def _emit_lines(chunk: ReplayableArrowChunk) -> AsyncIterator[str]:
259
+ blob = await loop.run_in_executor(None, _write_ndjson, chunk)
260
+ for line in blob.splitlines():
261
+ yield line.decode()
262
+
253
263
  async for item in heartbeat_over_stream(chunk_iter, total_timeout_s=total_timeout_s):
254
264
  if item is HEARTBEAT:
255
265
  yield HEARTBEAT
256
266
  continue
257
- pending[item.chunk_index] = item
267
+ pending.setdefault(item.chunk_index, []).append(item)
258
268
  while next_idx in pending:
259
- chunk = pending.pop(next_idx)
260
- blob = await loop.run_in_executor(None, _write_ndjson, chunk)
261
- for line in blob.splitlines():
262
- yield line.decode()
263
- next_idx += 1
269
+ chunk = pending[next_idx].pop(0)
270
+ if not pending[next_idx]:
271
+ del pending[next_idx]
272
+ next_idx += 1
273
+ async for line in _emit_lines(chunk):
274
+ yield line
275
+
276
+ # A genuine gap (an index that never arrived) must not strand chunks
277
+ # buffered past it -- drain whatever's left, in ascending index order.
278
+ for idx in sorted(pending):
279
+ for chunk in pending[idx]:
280
+ async for line in _emit_lines(chunk):
281
+ yield line
@@ -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
- Fast path: a single GET when already RUNNING, so this adds no
139
- meaningful overhead once warm."""
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
- async with httpx.AsyncClient() as client:
191
- await self._ensure_warehouse_running(client)
192
- resp = await self._authed_request(client, "POST", f"{self._host}/api/2.0/sql/statements", json=body)
193
- data = resp.json()
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
- _raise_for_failed(status)
204
- return data["statement_id"], data.get("manifest") or {}
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
- async with httpx.AsyncClient() as client:
253
- await self._authed_request(
254
- client,
255
- "PUT",
256
- f"{self._host}/api/2.0/fs/files{volume_path}",
257
- params={"overwrite": "true"},
258
- content_type="application/octet-stream",
259
- content=data,
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
- async with httpx.AsyncClient() as client:
267
- try:
268
- await self._authed_request(client, "DELETE", f"{self._host}/api/2.0/fs/files{volume_path}")
269
- except httpx.HTTPStatusError as exc:
270
- if exc.response.status_code != 404:
271
- raise
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
- async with httpx.AsyncClient() as client:
285
- async for blob, row_count, chunk_index in self._fetch_chunks_with_backpressure(
286
- client, statement_id, chunk_metas
287
- ):
288
- yield blob, row_count, chunk_index
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:
@@ -54,32 +54,57 @@ class _ResultSet:
54
54
  def __init__(self, schema: core.Schema | None, chunk_aiter: AsyncIterator[ReplayableArrowChunk]) -> None:
55
55
  self.schema = schema
56
56
  self._chunk_aiter = chunk_aiter
57
- self._pending: dict[int, ReplayableArrowChunk] = {}
57
+ # A list per index, not a single chunk -- a chunk_index can have more
58
+ # than one blob (DatabricksClient._fetch_chunk_index gathers over
59
+ # possibly-multiple `external_links` per chunk, "usually exactly one"
60
+ # but not guaranteed to be), and a dict keyed by index alone would
61
+ # silently overwrite/lose the first blob when a second one for the
62
+ # same index arrives.
63
+ self._pending: dict[int, list[ReplayableArrowChunk]] = {}
58
64
  self._next_idx = 0
59
65
  self._exhausted = False
60
66
  self._buffer: core.Table | None = None
61
67
  self.rownumber = 0
62
68
 
69
+ def _pop_pending(self, idx: int) -> ReplayableArrowChunk:
70
+ blobs = self._pending[idx]
71
+ chunk = blobs.pop(0)
72
+ if not blobs:
73
+ del self._pending[idx]
74
+ return chunk
75
+
63
76
  async def _pull_one_chunk_table(self) -> core.Table | None:
64
- """Returns the next expected chunk's Table, in order -- checking
65
- `_pending` FIRST, since a single earlier call can have pulled several
66
- chunks off `_chunk_aiter` before the one it actually needed showed up
67
- (arrival order is completion order, not chunk_index order), leaving
68
- the rest already-fetched-and-buffered here. Only touches the network
69
- (`_chunk_aiter.__anext__()`) once `_pending` has nothing more to give."""
77
+ """Returns one chunk's Table, preferring the next expected index --
78
+ checking `_pending` FIRST, since a single earlier call can have pulled
79
+ several chunks off `_chunk_aiter` before the one it actually needed
80
+ showed up (arrival order is completion order, not chunk_index order),
81
+ leaving the rest already-fetched-and-buffered here. Only touches the
82
+ network (`_chunk_aiter.__anext__()`) once `_pending` has nothing more
83
+ to give for the current index.
84
+
85
+ Once the source is exhausted, `_next_idx` reaching that exact value is
86
+ no longer required: if some index was skipped entirely (e.g. a chunk
87
+ whose bytes came back empty, see fetch_arrow_chunks_for_statement),
88
+ waiting for it forever would silently strand every higher-indexed
89
+ chunk already sitting in `_pending`. Draining the lowest remaining
90
+ index instead means a genuine gap costs row *order* past that point,
91
+ never lost rows."""
70
92
  while True:
71
93
  if self._next_idx in self._pending:
72
- ready = self._pending.pop(self._next_idx)
73
- self._next_idx += 1
74
- return ready.to_table()
94
+ chunk = self._pop_pending(self._next_idx)
95
+ if self._next_idx not in self._pending:
96
+ self._next_idx += 1
97
+ return chunk.to_table()
75
98
  if self._exhausted:
99
+ if self._pending:
100
+ return self._pop_pending(min(self._pending)).to_table()
76
101
  return None
77
102
  try:
78
103
  chunk = await self._chunk_aiter.__anext__()
79
104
  except StopAsyncIteration:
80
105
  self._exhausted = True
81
106
  continue
82
- self._pending[chunk.chunk_index] = chunk
107
+ self._pending.setdefault(chunk.chunk_index, []).append(chunk)
83
108
 
84
109
  async def _ensure_buffer(self, want: int) -> None:
85
110
  while (self._buffer is None or self._buffer.num_rows < want) and not self._exhausted:
@@ -233,6 +258,31 @@ class Cursor:
233
258
  async def fetchall_arrow(self) -> core.Table:
234
259
  return await self._require_result().fetchall_arrow()
235
260
 
261
+ def fetchall_streamed(self, *, total_timeout_s: float | None = None) -> AsyncIterator[Any]:
262
+ """Like fetchall(), but yields HEARTBEAT while pulling chunks instead
263
+ of blocking silently -- for a caller bridging e.g. an SSE connection
264
+ through the full download, not just the initial `execute_streamed`
265
+ wait for the statement to complete. Downloading many chunks for a
266
+ large result can itself take a while; `execute_streamed`'s own
267
+ heartbeats stop the moment the statement is ready, before any chunk
268
+ has actually been fetched. Yields HEARTBEAT zero or more times, then
269
+ the final `list[Row]`."""
270
+
271
+ async def _gen() -> AsyncIterator[Any]:
272
+ async for item in await_with_heartbeat(self.fetchall(), total_timeout_s=total_timeout_s):
273
+ yield item
274
+
275
+ return _gen()
276
+
277
+ def fetchall_arrow_streamed(self, *, total_timeout_s: float | None = None) -> AsyncIterator[Any]:
278
+ """Arrow-`Table` counterpart to fetchall_streamed -- see its docstring."""
279
+
280
+ async def _gen() -> AsyncIterator[Any]:
281
+ async for item in await_with_heartbeat(self.fetchall_arrow(), total_timeout_s=total_timeout_s):
282
+ yield item
283
+
284
+ return _gen()
285
+
236
286
  def __aiter__(self) -> Cursor:
237
287
  return self
238
288
 
@@ -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(router: respx.Router, n_chunks: int, rows_per_chunk: int, *, reverse_arrival: bool = False) -> None:
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
@@ -1,10 +1,12 @@
1
1
  from __future__ import annotations
2
2
 
3
+ import asyncio
4
+
3
5
  import httpx
4
6
  import pytest
5
7
  import respx
6
8
 
7
- from arrowbricks import HEARTBEAT, DatabricksClient
9
+ from arrowbricks import HEARTBEAT, DatabricksClient, QueryTimeout
8
10
  from arrowbricks.cursor import Cursor
9
11
 
10
12
 
@@ -153,6 +155,60 @@ async def test_execute_streamed_emits_heartbeat_before_cursor_ready(mock_warehou
153
155
  assert len(await cursor.fetchall()) == 3
154
156
 
155
157
 
158
+ @pytest.mark.asyncio
159
+ @respx.mock
160
+ async def test_fetchall_streamed_times_out_on_a_slow_chunk_download(warehouse_host_id):
161
+ """execute_streamed's own timeout only covers the wait for the statement
162
+ to become ready -- it stops the moment chunks are available to fetch,
163
+ before any chunk has actually been downloaded. A slow chunk download must
164
+ still be bounded, which is exactly what fetchall_streamed adds."""
165
+ host, warehouse_id = warehouse_host_id
166
+ respx.mock.get(f"{host}/api/2.0/sql/warehouses/{warehouse_id}").mock(
167
+ return_value=httpx.Response(200, json={"state": "RUNNING"})
168
+ )
169
+ respx.mock.post(f"{host}/api/2.0/sql/statements").mock(
170
+ return_value=httpx.Response(
171
+ 200,
172
+ json={
173
+ "statement_id": "stmt-slow",
174
+ "status": {"state": "SUCCEEDED"},
175
+ "manifest": {"chunks": [{"chunk_index": 0, "row_count": 3}]},
176
+ },
177
+ )
178
+ )
179
+ respx.mock.get(f"{host}/api/2.0/sql/statements/stmt-slow/result/chunks/0").mock(
180
+ return_value=httpx.Response(200, json={"external_links": [{"external_link": f"{host}/_data/slow-chunk"}]})
181
+ )
182
+
183
+ async def _slow_chunk(request: httpx.Request) -> httpx.Response:
184
+ await asyncio.sleep(10)
185
+ raise AssertionError("unreachable -- test should time out first")
186
+
187
+ respx.mock.get(f"{host}/_data/slow-chunk").mock(side_effect=_slow_chunk)
188
+ client = DatabricksClient(host, warehouse_id, token="test-token")
189
+ cursor = Cursor(client)
190
+ await cursor.execute("SELECT * FROM whatever")
191
+
192
+ with pytest.raises(QueryTimeout):
193
+ async for _ in cursor.fetchall_streamed(total_timeout_s=0.05):
194
+ pass
195
+
196
+
197
+ @pytest.mark.asyncio
198
+ @respx.mock
199
+ async def test_fetchall_streamed_yields_result_after_zero_or_more_heartbeats(mock_warehouse, warehouse_host_id):
200
+ host, warehouse_id = warehouse_host_id
201
+ mock_warehouse(respx.mock, n_chunks=1, rows_per_chunk=3)
202
+ client = DatabricksClient(host, warehouse_id, token="test-token")
203
+ cursor = Cursor(client)
204
+ await cursor.execute("SELECT * FROM whatever")
205
+
206
+ items = [item async for item in cursor.fetchall_streamed(total_timeout_s=5)]
207
+
208
+ assert all(item is HEARTBEAT for item in items[:-1])
209
+ assert [r[0] for r in items[-1]] == [0, 1, 2]
210
+
211
+
156
212
  @pytest.mark.asyncio
157
213
  async def test_fetch_before_execute_raises():
158
214
  client = DatabricksClient("https://fake", "wh", token="test-token")
@@ -0,0 +1,74 @@
1
+ """Regression tests for _ResultSet's chunk reorder buffer -- specifically the
2
+ two failure modes a naive dict[int, chunk] keyed-by-index buffer gets wrong:
3
+ a chunk_index that shows up more than once (DatabricksClient._fetch_chunk_index
4
+ gathers over possibly-multiple external_links per chunk), and a chunk_index
5
+ that never shows up at all (fetch_arrow_chunks_for_statement skips falsy
6
+ blobs). Both must never lose rows -- see cursor.py's _pull_one_chunk_table
7
+ docstring."""
8
+
9
+ from __future__ import annotations
10
+
11
+ from collections.abc import AsyncIterator
12
+
13
+ import pytest
14
+
15
+ from arrowbricks._streaming import ReplayableArrowChunk
16
+ from arrowbricks.cursor import _ResultSet
17
+
18
+
19
+ async def _aiter(chunks: list[ReplayableArrowChunk]) -> AsyncIterator[ReplayableArrowChunk]:
20
+ for c in chunks:
21
+ yield c
22
+
23
+
24
+ def _ids(chunk_bytes_builder, lo: int, hi: int, chunk_index: int) -> ReplayableArrowChunk:
25
+ return ReplayableArrowChunk(chunk_bytes_builder(lo, hi), chunk_index=chunk_index)
26
+
27
+
28
+ @pytest.mark.asyncio
29
+ async def test_duplicate_chunk_index_keeps_both_blobs(chunk_bytes_builder):
30
+ """Two separate blobs arriving under the SAME chunk_index (multiple
31
+ external_links for one chunk) must both survive, not overwrite each
32
+ other."""
33
+ chunks = [
34
+ _ids(chunk_bytes_builder, 0, 2, chunk_index=0),
35
+ _ids(chunk_bytes_builder, 5, 7, chunk_index=0), # same index, second blob
36
+ _ids(chunk_bytes_builder, 2, 4, chunk_index=1),
37
+ ]
38
+ result = _ResultSet(schema=None, chunk_aiter=_aiter(chunks))
39
+
40
+ table = await result.fetchall_arrow()
41
+
42
+ assert sorted(table.column(0).combine_chunks().to_pylist()) == [0, 1, 2, 3, 5, 6]
43
+
44
+
45
+ @pytest.mark.asyncio
46
+ async def test_gap_in_chunk_index_does_not_strand_later_chunks(chunk_bytes_builder):
47
+ """chunk_index 1 never arrives at all (e.g. its bytes came back empty and
48
+ got filtered upstream) -- chunk 2's rows must still come out, not get
49
+ stranded behind a hole that never fills in."""
50
+ chunks = [
51
+ _ids(chunk_bytes_builder, 0, 2, chunk_index=0),
52
+ _ids(chunk_bytes_builder, 2, 4, chunk_index=2), # index 1 is missing
53
+ ]
54
+ result = _ResultSet(schema=None, chunk_aiter=_aiter(chunks))
55
+
56
+ table = await result.fetchall_arrow()
57
+
58
+ assert sorted(table.column(0).combine_chunks().to_pylist()) == [0, 1, 2, 3]
59
+
60
+
61
+ @pytest.mark.asyncio
62
+ async def test_out_of_order_arrival_still_preserves_row_order(chunk_bytes_builder):
63
+ """The common case (no gaps, no duplicates, just network completion
64
+ order != chunk_index order) must still come out in chunk_index order."""
65
+ chunks = [
66
+ _ids(chunk_bytes_builder, 4, 6, chunk_index=2),
67
+ _ids(chunk_bytes_builder, 0, 2, chunk_index=0),
68
+ _ids(chunk_bytes_builder, 2, 4, chunk_index=1),
69
+ ]
70
+ result = _ResultSet(schema=None, chunk_aiter=_aiter(chunks))
71
+
72
+ table = await result.fetchall_arrow()
73
+
74
+ assert table.column(0).combine_chunks().to_pylist() == [0, 1, 2, 3, 4, 5]
@@ -0,0 +1,82 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+
5
+ import httpx
6
+ import pytest
7
+ import respx
8
+
9
+ from arrowbricks import DatabricksClient, stream_query_json
10
+
11
+
12
+ @pytest.mark.asyncio
13
+ @respx.mock
14
+ async def test_stream_query_json_preserves_order_despite_out_of_order_chunks(mock_warehouse, warehouse_host_id):
15
+ host, warehouse_id = warehouse_host_id
16
+ # Chunks resolve in REVERSE completion order (see conftest's
17
+ # reverse_arrival) -- this is the scenario that silently breaks ORDER BY
18
+ # if a caller doesn't reorder before emitting.
19
+ mock_warehouse(respx.mock, n_chunks=4, rows_per_chunk=5, reverse_arrival=True)
20
+ client = DatabricksClient(host, warehouse_id, token="test-token")
21
+
22
+ rows = [json.loads(row) async for row in stream_query_json(client, "SELECT * FROM whatever ORDER BY id")]
23
+
24
+ assert [r["id"] for r in rows] == list(range(20))
25
+
26
+
27
+ @pytest.mark.asyncio
28
+ @respx.mock
29
+ async def test_stream_query_json_row_shape(mock_warehouse, warehouse_host_id):
30
+ host, warehouse_id = warehouse_host_id
31
+ mock_warehouse(respx.mock, n_chunks=1, rows_per_chunk=3)
32
+ client = DatabricksClient(host, warehouse_id, token="test-token")
33
+
34
+ rows = [json.loads(row) async for row in stream_query_json(client, "SELECT * FROM whatever")]
35
+
36
+ assert rows == [
37
+ {"id": 0, "label": "row_0"},
38
+ {"id": 1, "label": "row_1"},
39
+ {"id": 2, "label": "row_2"},
40
+ ]
41
+
42
+
43
+ @pytest.mark.asyncio
44
+ @respx.mock
45
+ async def test_stream_query_json_survives_duplicate_index_and_a_gap(warehouse_host_id, chunk_bytes_builder):
46
+ """chunk_index 0 has TWO external_links (a real, if uncommon, shape --
47
+ DatabricksClient._fetch_chunk_index gathers over all of them), and index 1
48
+ never shows up at all. Neither must lose rows -- see cursor.py's
49
+ _pull_one_chunk_table / this module's stream_query_json docstrings."""
50
+ host, warehouse_id = warehouse_host_id
51
+ statement_id = "stmt-dup-gap"
52
+ respx.mock.get(f"{host}/api/2.0/sql/warehouses/{warehouse_id}").mock(
53
+ return_value=httpx.Response(200, json={"state": "RUNNING"})
54
+ )
55
+ respx.mock.post(f"{host}/api/2.0/sql/statements").mock(
56
+ return_value=httpx.Response(
57
+ 200,
58
+ json={
59
+ "statement_id": statement_id,
60
+ "status": {"state": "SUCCEEDED"},
61
+ # index 1 deliberately absent -- a genuine gap.
62
+ "manifest": {"chunks": [{"chunk_index": 0, "row_count": 4}, {"chunk_index": 2, "row_count": 2}]},
63
+ },
64
+ )
65
+ )
66
+ respx.mock.get(f"{host}/api/2.0/sql/statements/{statement_id}/result/chunks/0").mock(
67
+ return_value=httpx.Response(
68
+ 200,
69
+ json={"external_links": [{"external_link": f"{host}/_data/0a"}, {"external_link": f"{host}/_data/0b"}]},
70
+ )
71
+ )
72
+ respx.mock.get(f"{host}/api/2.0/sql/statements/{statement_id}/result/chunks/2").mock(
73
+ return_value=httpx.Response(200, json={"external_links": [{"external_link": f"{host}/_data/2"}]})
74
+ )
75
+ respx.mock.get(f"{host}/_data/0a").mock(return_value=httpx.Response(200, content=chunk_bytes_builder(0, 2)))
76
+ respx.mock.get(f"{host}/_data/0b").mock(return_value=httpx.Response(200, content=chunk_bytes_builder(10, 12)))
77
+ respx.mock.get(f"{host}/_data/2").mock(return_value=httpx.Response(200, content=chunk_bytes_builder(20, 22)))
78
+ client = DatabricksClient(host, warehouse_id, token="test-token")
79
+
80
+ rows = [json.loads(row) async for row in stream_query_json(client, "SELECT * FROM whatever")]
81
+
82
+ assert sorted(r["id"] for r in rows) == [0, 1, 10, 11, 20, 21]
@@ -145,7 +145,7 @@ wheels = [
145
145
 
146
146
  [[package]]
147
147
  name = "arrowbricks"
148
- version = "0.1.2"
148
+ version = "0.3.0"
149
149
  source = { editable = "." }
150
150
  dependencies = [
151
151
  { name = "arro3-core" },
@@ -1,39 +0,0 @@
1
- from __future__ import annotations
2
-
3
- import json
4
-
5
- import pytest
6
- import respx
7
-
8
- from arrowbricks import DatabricksClient, stream_query_json
9
-
10
-
11
- @pytest.mark.asyncio
12
- @respx.mock
13
- async def test_stream_query_json_preserves_order_despite_out_of_order_chunks(mock_warehouse, warehouse_host_id):
14
- host, warehouse_id = warehouse_host_id
15
- # Chunks resolve in REVERSE completion order (see conftest's
16
- # reverse_arrival) -- this is the scenario that silently breaks ORDER BY
17
- # if a caller doesn't reorder before emitting.
18
- mock_warehouse(respx.mock, n_chunks=4, rows_per_chunk=5, reverse_arrival=True)
19
- client = DatabricksClient(host, warehouse_id, token="test-token")
20
-
21
- rows = [json.loads(row) async for row in stream_query_json(client, "SELECT * FROM whatever ORDER BY id")]
22
-
23
- assert [r["id"] for r in rows] == list(range(20))
24
-
25
-
26
- @pytest.mark.asyncio
27
- @respx.mock
28
- async def test_stream_query_json_row_shape(mock_warehouse, warehouse_host_id):
29
- host, warehouse_id = warehouse_host_id
30
- mock_warehouse(respx.mock, n_chunks=1, rows_per_chunk=3)
31
- client = DatabricksClient(host, warehouse_id, token="test-token")
32
-
33
- rows = [json.loads(row) async for row in stream_query_json(client, "SELECT * FROM whatever")]
34
-
35
- assert rows == [
36
- {"id": 0, "label": "row_0"},
37
- {"id": 1, "label": "row_1"},
38
- {"id": 2, "label": "row_2"},
39
- ]
File without changes
File without changes
File without changes