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.
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/AGENTS.md +4 -3
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/PKG-INFO +8 -4
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/README.md +7 -3
- arrowbricks-0.3.0/examples/fastapi_sse_pivot.py +78 -0
- {arrowbricks-0.1.2 → 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.1.2 → arrowbricks-0.3.0}/src/arrowbricks/_streaming.py +26 -8
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/src/arrowbricks/client.py +81 -34
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/src/arrowbricks/cursor.py +61 -11
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/tests/conftest.py +5 -2
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/tests/test_client.py +81 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/tests/test_cursor.py +57 -1
- arrowbricks-0.3.0/tests/test_result_set.py +74 -0
- arrowbricks-0.3.0/tests/test_streaming.py +82 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/uv.lock +1 -1
- arrowbricks-0.1.2/tests/test_streaming.py +0 -39
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/.github/workflows/ci.yml +0 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/.github/workflows/release.yml +0 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/.gitignore +0 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/LICENSE +0 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/examples/azure_auth.py +0 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/examples/basic.py +0 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/examples/cursor_paging.py +0 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/examples/fastapi_sse.py +0 -0
- {arrowbricks-0.1.2 → arrowbricks-0.3.0}/prek.toml +0 -0
- {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
|
|
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.
|
|
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,
|
|
76
|
-
for a
|
|
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,
|
|
63
|
-
for a
|
|
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.
|
|
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
|
|
267
|
+
pending.setdefault(item.chunk_index, []).append(item)
|
|
258
268
|
while next_idx in pending:
|
|
259
|
-
chunk = pending.pop(
|
|
260
|
-
|
|
261
|
-
|
|
262
|
-
|
|
263
|
-
|
|
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
|
-
|
|
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:
|
|
@@ -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
|
-
|
|
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
|
|
65
|
-
`_pending` FIRST, since a single earlier call can have pulled
|
|
66
|
-
chunks off `_chunk_aiter` before the one it actually needed
|
|
67
|
-
(arrival order is completion order, not chunk_index order),
|
|
68
|
-
the rest already-fetched-and-buffered here. Only touches the
|
|
69
|
-
(`_chunk_aiter.__anext__()`) once `_pending` has nothing more
|
|
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
|
-
|
|
73
|
-
self._next_idx
|
|
74
|
-
|
|
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
|
|
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(
|
|
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]
|
|
@@ -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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|