arrowbricks 1.3.2__tar.gz → 1.3.3__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 (30) hide show
  1. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/PKG-INFO +2 -2
  2. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/README.md +1 -1
  3. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/pyproject.toml +3 -3
  4. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/Cargo.lock +1 -1
  5. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/Cargo.toml +1 -1
  6. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/README.md +6 -7
  7. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/examples/duckdb_query.py +2 -1
  8. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/src/client.rs +79 -12
  9. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/src/lib.rs +42 -151
  10. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/src/pipeline.rs +113 -0
  11. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/tests/wiremock_pipeline.rs +185 -0
  12. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/tests_py/test_parameters.py +3 -16
  13. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/tests_py/test_streaming.py +42 -33
  14. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/tests_py/test_token_provider.py +52 -2
  15. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/src/arrowbricks/_core.pyi +0 -27
  16. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/src/arrowbricks/client.py +18 -18
  17. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/src/arrowbricks/cursor.py +13 -0
  18. arrowbricks-1.3.2/rust/arrowbricks_core/tests_py/test_execute_json.py +0 -101
  19. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/LICENSE +0 -0
  20. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/.gitignore +0 -0
  21. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/examples/fastapi_sse.py +0 -0
  22. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/rustfmt.toml +0 -0
  23. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/src/heartbeat.rs +0 -0
  24. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/tests/wiremock_volume_files.rs +0 -0
  25. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/tests_py/test_ipc_stream.py +0 -0
  26. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/tests_py/test_stream_ndjson_lines.py +0 -0
  27. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/rust/arrowbricks_core/tests_py/test_volume_files.py +0 -0
  28. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/src/arrowbricks/__init__.py +0 -0
  29. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/src/arrowbricks/_streaming.py +0 -0
  30. {arrowbricks-1.3.2 → arrowbricks-1.3.3}/src/arrowbricks/py.typed +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: arrowbricks
3
- Version: 1.3.2
3
+ Version: 1.3.3
4
4
  Requires-Dist: arro3-core>=0.8 ; extra == 'arro3'
5
5
  Provides-Extra: arro3
6
6
  License-File: LICENSE
@@ -191,7 +191,7 @@ con.register("my_table", table)
191
191
  con.sql("SELECT count(*) FROM my_table").show()
192
192
  ```
193
193
 
194
- `ReplayableArrowChunk` works the same way for Arrow-IPC bytes you fetched and stored earlier (e.g. `client.execute_arrow_statement(...)`'s raw chunk bytes, cached in Redis/a file/wherever) -- DuckDB's registration path calls `__arrow_c_stream__` twice (a schema peek, then the actual scan), which is exactly what `ReplayableArrowChunk` exists to support:
194
+ `ReplayableArrowChunk` works the same way for Arrow-IPC bytes you fetched and stored earlier (e.g. a raw chunk's bytes, cached in Redis/a file/wherever) -- DuckDB's registration path calls `__arrow_c_stream__` twice (a schema peek, then the actual scan), which is exactly what `ReplayableArrowChunk` exists to support:
195
195
 
196
196
  ```python
197
197
  from arrowbricks import ReplayableArrowChunk
@@ -179,7 +179,7 @@ con.register("my_table", table)
179
179
  con.sql("SELECT count(*) FROM my_table").show()
180
180
  ```
181
181
 
182
- `ReplayableArrowChunk` works the same way for Arrow-IPC bytes you fetched and stored earlier (e.g. `client.execute_arrow_statement(...)`'s raw chunk bytes, cached in Redis/a file/wherever) -- DuckDB's registration path calls `__arrow_c_stream__` twice (a schema peek, then the actual scan), which is exactly what `ReplayableArrowChunk` exists to support:
182
+ `ReplayableArrowChunk` works the same way for Arrow-IPC bytes you fetched and stored earlier (e.g. a raw chunk's bytes, cached in Redis/a file/wherever) -- DuckDB's registration path calls `__arrow_c_stream__` twice (a schema peek, then the actual scan), which is exactly what `ReplayableArrowChunk` exists to support:
183
183
 
184
184
  ```python
185
185
  from arrowbricks import ReplayableArrowChunk
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "arrowbricks"
3
- version = "1.3.2"
3
+ version = "1.3.3"
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. Rust/PyO3 core throughout -- zero required runtime dependencies."
5
5
  readme = "README.md"
6
6
  license = "MIT"
@@ -10,8 +10,8 @@ requires-python = ">=3.11"
10
10
  # (see rust/arrowbricks_core's own reqwest client + Arrow-IPC/NDJSON
11
11
  # writer). arro3-core is optional (see below): only Cursor.fetchone/
12
12
  # fetchmany/fetchall's row materialization needs it -- fetchall_arrow,
13
- # execute_arrow, stream_query_json, upload_volume_file/delete_volume_file,
14
- # and Cursor.description all work with zero dependencies installed.
13
+ # stream_query_json, upload_volume_file/delete_volume_file, and
14
+ # Cursor.description all work with zero dependencies installed.
15
15
  dependencies = []
16
16
 
17
17
  [project.optional-dependencies]
@@ -240,7 +240,7 @@ dependencies = [
240
240
 
241
241
  [[package]]
242
242
  name = "arrowbricks_core"
243
- version = "1.3.2"
243
+ version = "1.3.3"
244
244
  dependencies = [
245
245
  "arrow",
246
246
  "arrow-json",
@@ -1,6 +1,6 @@
1
1
  [package]
2
2
  name = "arrowbricks_core"
3
- version = "1.3.2"
3
+ version = "1.3.3"
4
4
  edition = "2024"
5
5
  readme = "README.md"
6
6
 
@@ -48,7 +48,8 @@ async def main():
48
48
  warehouse_id="abcd1234efgh5678",
49
49
  token="dapi...",
50
50
  )
51
- table = await client.execute_arrow("SELECT * FROM my_catalog.my_schema.my_table LIMIT 100")
51
+ result = await client.execute("SELECT * FROM my_catalog.my_schema.my_table LIMIT 100")
52
+ table = await result.fetchall_arrow()
52
53
  print(table.num_rows)
53
54
 
54
55
 
@@ -66,17 +67,14 @@ unlike the table object itself.
66
67
  ## API
67
68
 
68
69
  - `Client(host, warehouse_id, *, token=None, token_provider=None, chunk_fetch_concurrency=64, http_timeout=60.0, wait_timeout="30s", warehouse_start_timeout=300.0, warehouse_confirmed_running_ttl_s=30.0, compress_results=True)` -- exactly one of `token`/`token_provider`. `token_provider` is a callable (sync or async) returning a token string, called fresh on every request, no caching. `compress_results` requests LZ4-compressed cloud-fetch chunks (see above); set `False` to opt out.
69
- - `Client.execute_arrow(statement, *, catalog=None, schema=None, parameters=None) -> Table` -- eager: fetches and assembles the whole result before returning. `parameters` is Databricks' own named-parameter format (`[{"name":..., "value":..., "type":...}]`), passed straight through.
70
- - `Client.execute(statement, *, catalog=None, schema=None, parameters=None) -> ResultSet` -- submits and starts background chunk fetching without pulling anything yet.
70
+ - `Client.execute(statement, *, catalog=None, schema=None, parameters=None) -> ResultSet` -- submits and starts background chunk fetching without pulling anything yet. `parameters` is Databricks' own named-parameter format (`[{"name":..., "value":..., "type":...}]`), passed straight through.
71
71
  - `ResultSet.fetchmany_arrow(n) -> Table` -- pulls/decodes only as many chunks as needed for `n` rows, buffering the rest; may return fewer than `n` once exhausted.
72
72
  - `ResultSet.fetchall_arrow() -> Table` -- drains everything remaining.
73
+ - `ResultSet.fetchall_arrow_streamed(*, total_timeout_s=None)` -- same as `fetchall_arrow()`, but an async iterator yielding the `HEARTBEAT` singleton while pulling chunks instead of blocking silently (bridge e.g. an SSE connection through the download), then a `Table` exactly once. Raises if `total_timeout_s` elapses first.
73
74
  - `ResultSet.schema() -> list[tuple[str, str]] | None` -- the real decoded Arrow schema as `(name, type_name)` string pairs, once known (after >=1 fetch); `None` before that. Computed directly in Rust -- unlike a `Table`'s own `.schema`, this needs no arro3 install.
74
75
  - `ResultSet.statement_id`, `ResultSet.num_chunks`, `ResultSet.columns` (manifest-based pre-fetch schema estimate, also arro3-free).
75
- - `Client.execute_json(statement, *, catalog=None, schema=None, parameters=None) -> list[list[str | None]]` -- JSON_ARRAY format, no Arrow parse. Every non-null value comes back as a string regardless of real column type (Databricks' own contract) -- cast by the manifest's column type yourself if you want native Python types.
76
76
  - `Client.stream_ndjson_lines(statement, *, catalog=None, schema=None, parameters=None, total_timeout_s=None)` -- chunk-at-a-time: an async iterator yielding the `HEARTBEAT` singleton while waiting on the statement or any individual chunk, then a `list[str]` of NDJSON lines (one per row, explicit nulls, ISO-8601 timestamps) per chunk in logical order. Decode and JSON encoding both happen in Rust -- backs `stream_query_json` end to end.
77
77
  - `Client.upload_volume_file(volume_path, data: bytes)` / `Client.delete_volume_file(volume_path)` -- Unity Catalog volume files via the Files API. Delete treats a 404 as success (idempotent). Both raise a plain `RuntimeError` (message only) on failure.
78
- - `Client.execute_streamed(statement, *, catalog=None, schema=None, parameters=None, total_timeout_s=None)` -- like `execute()`, but an async iterator yielding the `HEARTBEAT` singleton while waiting on Databricks instead of blocking silently (bridge e.g. an SSE connection through a slow warehouse cold start), then a `ResultSet` exactly once. Raises if `total_timeout_s` elapses first.
79
- - `ResultSet.fetchall_arrow_streamed(*, total_timeout_s=None)` -- same idea for the chunk-download phase: yields `HEARTBEAT` while pulling chunks, then a `Table` exactly once.
80
78
  - `write_ipc_stream(stream, buf)` -- free function; writes any object implementing `__arrow_c_stream__` (a `Table` from this crate, arro3, pyarrow, ...) as uncompressed Arrow-IPC stream bytes to a Python file-like object. No dependency needed regardless of the input's origin.
81
79
  - `read_ipc_stream(data: bytes) -> Table` -- free function; the exact inverse of `write_ipc_stream`, parsing raw Arrow-IPC stream bytes back into a `Table`. No dependency needed regardless of where the bytes came from -- backs `arrowbricks.ReplayableArrowChunk`, which needs to re-parse the same cached bytes on every `__arrow_c_stream__` call.
82
80
  - `HEARTBEAT` -- module-level singleton; compare with `is`, e.g. `if item is _core.HEARTBEAT: ...`.
@@ -90,7 +88,8 @@ directly, no pandas/pyarrow conversion step, no extra dependency:
90
88
  ```python
91
89
  import duckdb
92
90
 
93
- result = await client.execute_arrow("SELECT * FROM my_catalog.my_schema.my_table")
91
+ result_set = await client.execute("SELECT * FROM my_catalog.my_schema.my_table")
92
+ result = await result_set.fetchall_arrow()
94
93
  print(duckdb.sql("SELECT count(*) FROM result").fetchall())
95
94
  ```
96
95
 
@@ -27,7 +27,8 @@ async def main() -> None:
27
27
  warehouse_id=os.environ["DATABRICKS_WAREHOUSE_ID"],
28
28
  token=os.environ["DATABRICKS_TOKEN"],
29
29
  )
30
- result = await client.execute_arrow("SELECT * FROM my_catalog.my_schema.my_table LIMIT 1000000") # noqa: F841 -- read by DuckDB's replacement scan via local variable name, not a literal Python reference
30
+ result_set = await client.execute("SELECT * FROM my_catalog.my_schema.my_table LIMIT 1000000")
31
+ result = await result_set.fetchall_arrow() # noqa: F841 -- read by DuckDB's replacement scan via local variable name, not a literal Python reference
31
32
 
32
33
  # `result` is queryable by variable name -- DuckDB imports it zero-copy,
33
34
  # no pandas/pyarrow conversion step in between.
@@ -197,7 +197,18 @@ impl ApiError {
197
197
  /// (excluding `is_connect()`/`is_timeout()`, already handled) also covers
198
198
  /// hyper's own "connection closed before message completed" pooled-
199
199
  /// connection-reuse race, which fires before any body is even read and
200
- /// is equally safe to retry on an idempotent request.
200
+ /// is equally safe to retry on an idempotent request. Unlike `is_decode()`
201
+ /// (hit for real on a 400-chunk fetch and reproduced with a directed
202
+ /// test), this specific race is reasoned from reqwest/hyper's own source,
203
+ /// not independently reproduced -- a directed attempt to force it
204
+ /// self-healed 20/20 times (hyper's own idle-connection health check
205
+ /// evidently detects a closed peer and opens a fresh connection before
206
+ /// ever handing the dead one out for reuse, at least under a simple,
207
+ /// low-concurrency test). Kept anyway since it's safe regardless (scoped
208
+ /// to idempotent requests only) and the window may still be reachable
209
+ /// under real production concurrency even though a quick local test
210
+ /// couldn't force it -- treat as a plausible defensive measure, not a
211
+ /// confirmed fix for an observed failure.
201
212
  fn from_reqwest(e: reqwest::Error, idempotent: bool) -> Self {
202
213
  // `is_decode()` (reqwest's `Kind::Decode`) isn't only content-decoding
203
214
  // -- `Response::bytes()`/`.text()`/`.json()` also wrap a body read
@@ -219,8 +230,18 @@ impl ApiError {
219
230
  }
220
231
  }
221
232
 
222
- fn from_status(status: StatusCode, body: &str) -> Self {
223
- let transient = matches!(status.as_u16(), 401 | 403 | 408 | 429) || status.is_server_error();
233
+ /// `idempotent` only gates the 5xx case -- 401/403/408/429 mean the
234
+ /// request was rejected before any processing started (auth failure,
235
+ /// rate limit, client-side timeout), safe to retry regardless of method.
236
+ /// A 5xx is murkier: it usually means the same, but it can also mean the
237
+ /// backend already accepted and started the statement before some later
238
+ /// failure (e.g. a gateway timeout) produced the 5xx anyway -- found in
239
+ /// code review that this call was unconditionally transient even for the
240
+ /// statement-submit POST, bypassing the exact idempotency reasoning
241
+ /// `from_reqwest`'s `idempotent` param exists for (retrying that POST
242
+ /// risks a second, duplicate execution of arbitrary caller SQL).
243
+ fn from_status(status: StatusCode, body: &str, idempotent: bool) -> Self {
244
+ let transient = matches!(status.as_u16(), 401 | 403 | 408 | 429) || (idempotent && status.is_server_error());
224
245
  Self {
225
246
  message: format!("HTTP {status}: {body}"),
226
247
  transient,
@@ -316,14 +337,22 @@ fn decompress_lz4_frame(compressed: &Bytes) -> Result<Bytes, ApiError> {
316
337
  // small lower bound.
317
338
  let mut out = Vec::with_capacity(compressed.len() * 4);
318
339
  let mut decoder = lz4_flex::frame::FrameDecoder::new(&compressed[..]);
319
- loop {
320
- let before = out.len();
340
+ // Terminate on the *reader* being exhausted, not on "output stopped
341
+ // growing" -- found in code review that a frame which happens to decode
342
+ // to zero bytes (a real, valid LZ4 Frame shape: header + immediate
343
+ // EndMark) makes one `read_to_end` call return `Ok(0)` for that frame
344
+ // without erroring and without necessarily advancing into the next
345
+ // frame yet, which the old `out.len() == before` check read as "no more
346
+ // frames" -- silently dropping every subsequent concatenated frame with
347
+ // no error at all. Same silent-truncation shape as the original
348
+ // multi-frame bug this loop exists to fix. Looping while the reader
349
+ // still has bytes left (regardless of whether the last call grew `out`)
350
+ // is the correct fix -- verified against a zero-content frame sandwiched
351
+ // between two real ones, see `decompress_lz4_frame_survives_a_zero_content_frame_in_the_middle`.
352
+ while !decoder.get_ref().is_empty() {
321
353
  decoder
322
354
  .read_to_end(&mut out)
323
355
  .map_err(|e| ApiError::permanent(format!("LZ4 frame decompress failed: {e}")))?;
324
- if out.len() == before {
325
- break;
326
- }
327
356
  }
328
357
  Ok(Bytes::from(out))
329
358
  }
@@ -499,7 +528,7 @@ impl DbClient {
499
528
  let status = resp.status();
500
529
  let text = resp.text().await.map_err(|e| ApiError::from_reqwest(e, idempotent))?;
501
530
  if !status.is_success() {
502
- return Err(ApiError::from_status(status, &text));
531
+ return Err(ApiError::from_status(status, &text, idempotent));
503
532
  }
504
533
  serde_json::from_str::<T>(&text).map_err(|e| ApiError::permanent(format!("bad JSON body: {e}")))
505
534
  })
@@ -530,7 +559,7 @@ impl DbClient {
530
559
  let status = resp.status();
531
560
  if !status.is_success() {
532
561
  let text = resp.text().await.unwrap_or_default();
533
- return Err(ApiError::from_status(status, &text));
562
+ return Err(ApiError::from_status(status, &text, true));
534
563
  }
535
564
  let bytes = resp.bytes().await.map_err(|e| ApiError::from_reqwest(e, true))?;
536
565
  if !compressed {
@@ -837,7 +866,7 @@ impl DbClient {
837
866
  let status = resp.status();
838
867
  if !status.is_success() {
839
868
  let text = resp.text().await.unwrap_or_default();
840
- return Err(ApiError::from_status(status, &text));
869
+ return Err(ApiError::from_status(status, &text, true));
841
870
  }
842
871
  Ok(())
843
872
  })
@@ -866,7 +895,7 @@ impl DbClient {
866
895
  return Ok(());
867
896
  }
868
897
  let text = resp.text().await.unwrap_or_default();
869
- Err(ApiError::from_status(status, &text))
898
+ Err(ApiError::from_status(status, &text, true))
870
899
  })
871
900
  .await
872
901
  }
@@ -932,6 +961,44 @@ mod tests {
932
961
  assert_eq!(decompressed, Bytes::from(expected));
933
962
  }
934
963
 
964
+ /// Regression test for a bug found in code review: a real, valid LZ4
965
+ /// Frame that happens to decode to zero bytes (a header immediately
966
+ /// followed by an EndMark -- a legal frame shape, not malformed input)
967
+ /// makes `read_to_end` return `Ok(0)` for that frame without erroring.
968
+ /// The old loop read "output didn't grow" as "no more frames" and
969
+ /// stopped there, silently dropping every frame concatenated after the
970
+ /// empty one. `decompress_lz4_frame` must keep going as long as the
971
+ /// underlying reader still has bytes left, not just as long as output
972
+ /// keeps growing.
973
+ #[test]
974
+ fn decompress_lz4_frame_survives_a_zero_content_frame_in_the_middle() {
975
+ use std::io::Write;
976
+
977
+ fn compress_one_frame(data: &[u8]) -> Vec<u8> {
978
+ let mut encoder = lz4_flex::frame::FrameEncoder::new(Vec::new());
979
+ encoder.write_all(data).unwrap();
980
+ encoder.finish().unwrap()
981
+ }
982
+
983
+ let part_a = b"the quick brown fox jumps over the lazy dog ".repeat(50);
984
+ let part_b = b"pack my box with five dozen liquor jugs ".repeat(50);
985
+ let mut concatenated_frames = Vec::new();
986
+ concatenated_frames.extend(compress_one_frame(&part_a));
987
+ concatenated_frames.extend(compress_one_frame(b"")); // real frame, zero content
988
+ concatenated_frames.extend(compress_one_frame(&part_b));
989
+
990
+ let decompressed = decompress_lz4_frame(&Bytes::from(concatenated_frames)).unwrap();
991
+
992
+ let mut expected = Vec::new();
993
+ expected.extend_from_slice(&part_a);
994
+ expected.extend_from_slice(&part_b);
995
+ assert_eq!(
996
+ decompressed,
997
+ Bytes::from(expected),
998
+ "the frame after the zero-content one must not be silently dropped"
999
+ );
1000
+ }
1001
+
935
1002
  fn ok_task() -> tokio::task::JoinHandle<Result<(), ApiError>> {
936
1003
  tokio::spawn(async { Ok(()) })
937
1004
  }
@@ -19,11 +19,6 @@ use client::{ApiError, DbClient, TokenFuture, TokenProvider};
19
19
  use heartbeat::{HeartbeatStream, HeartbeatWait, Tick};
20
20
  use pipeline::{NdjsonStream, ResultStream};
21
21
 
22
- #[pyfunction]
23
- fn ping() -> &'static str {
24
- "pong"
25
- }
26
-
27
22
  /// Writes any object implementing `__arrow_c_stream__` (a `Table`/
28
23
  /// `RecordBatchReader` from this crate, arro3, pyarrow, or anything else
29
24
  /// Arrow-C-Data-Interface-compatible) as Arrow-IPC stream bytes to `buf` (a
@@ -125,19 +120,32 @@ fn heartbeat_singleton(py: Python<'_>) -> PyResult<Py<PyHeartbeat>> {
125
120
  /// whatever this future awaits.
126
121
  ///
127
122
  /// `get_token` isn't only ever called from the outermost task
128
- /// `future_into_py` wraps -- `execute()`/`execute_arrow()` spawn chunk-fetch
123
+ /// `future_into_py` wraps -- `execute()`'s `ResultSet` spawns chunk-fetch
129
124
  /// worker tasks (`fetch_chunks_with_backpressure`) that each call an
130
125
  /// authenticated endpoint (chunk-index resolution) too, and those inner
131
126
  /// `tokio::spawn`ed tasks don't inherit the outer task's asyncio-event-loop
132
127
  /// context. Calling `pyo3_async_runtimes::tokio::into_future` from one of
133
128
  /// them fails with "no running event loop" -- caught by testing an async
134
129
  /// `token_provider` against a multi-chunk result, not by the eager/sync
135
- /// cases alone. Fix: capture the current task's `TaskLocals` on first use
136
- /// (guaranteed to be the outer context, since `execute_arrow_statement`
137
- /// always needs a token before any worker is spawned) and cache it, so
138
- /// later calls -- including from worker tasks -- run inside
139
- /// `pyo3_async_runtimes::tokio::scope` with that same captured context
140
- /// instead of trying to discover one from whatever task happens to call.
130
+ /// cases alone. Fix: capture the current task's `TaskLocals` (guaranteed to
131
+ /// be the outer context, since `execute_arrow_statement` always needs a
132
+ /// token before any worker is spawned) and cache it, so later calls --
133
+ /// including from worker tasks -- run inside `pyo3_async_runtimes::tokio::scope`
134
+ /// with that same captured context instead of trying to discover one from
135
+ /// whatever task happens to call.
136
+ ///
137
+ /// The capture re-runs on **every** call from a context that has its own
138
+ /// loop, not just the first ever -- found in code review that caching only
139
+ /// once-if-`None` pins the `DbClient` (which persists across many separate
140
+ /// `execute()` calls, by design -- see its own doc comment) to whichever
141
+ /// event loop happened to be running the very first time a token was ever
142
+ /// requested. A second, later `asyncio.run()` (or any fresh loop -- a
143
+ /// restarted worker, a new pytest-asyncio test) then scopes the awaited
144
+ /// provider onto a loop that's already closed, failing with "Event loop is
145
+ /// closed" instead of just using the current one. Re-capturing costs
146
+ /// nothing extra from a worker task (its own attempt fails, same as before,
147
+ /// falling through to whatever's cached) since the outer call for *this*
148
+ /// statement already refreshed the cache before any worker was spawned.
141
149
  struct PyTokenProvider {
142
150
  callable: Py<PyAny>,
143
151
  locals: Mutex<Option<TaskLocals>>,
@@ -148,13 +156,15 @@ impl TokenProvider for PyTokenProvider {
148
156
  let callable = Python::attach(|py| self.callable.clone_ref(py));
149
157
  let locals = {
150
158
  let mut guard = self.locals.lock().unwrap();
151
- if guard.is_none() {
152
- // Best-effort: if this call isn't in a context with a
153
- // running loop either, leave it None and fall through to
154
- // the no-scope path below (matches today's behavior).
155
- if let Ok(captured) = Python::attach(pyo3_async_runtimes::tokio::get_current_locals) {
156
- *guard = Some(captured);
157
- }
159
+ // Unconditional, not `if guard.is_none()` -- see this struct's
160
+ // own doc comment for why caching only once pinned the client to
161
+ // its first-ever event loop. Best-effort: if this particular
162
+ // call isn't in a context with a running loop either (a worker
163
+ // task), leave whatever's already cached alone and fall through
164
+ // to using that (or the no-scope path below, if nothing has ever
165
+ // been captured at all).
166
+ if let Ok(captured) = Python::attach(pyo3_async_runtimes::tokio::get_current_locals) {
167
+ *guard = Some(captured);
158
168
  }
159
169
  guard.clone()
160
170
  };
@@ -192,8 +202,8 @@ impl TokenProvider for PyTokenProvider {
192
202
  }
193
203
 
194
204
  /// One Databricks SQL warehouse endpoint -- a persistent `reqwest::Client`
195
- /// (connection pool) reused across every `execute`/`execute_arrow` call,
196
- /// same reason Python's own `DatabricksClient` reuses one `httpx.AsyncClient`
205
+ /// (connection pool) reused across every `execute` call, same reason
206
+ /// Python's own `DatabricksClient` reuses one `httpx.AsyncClient`
197
207
  /// rather than building a fresh one per statement: repeated TCP+TLS
198
208
  /// handshakes are pure waste against the same host.
199
209
  #[pyclass(name = "Client")]
@@ -234,7 +244,12 @@ impl PyDbClient {
234
244
  compress_results: bool,
235
245
  ) -> PyResult<Self> {
236
246
  let db_client = match (token, token_provider) {
237
- (Some(t), _) => DbClient::new(&host, &warehouse_id, &t),
247
+ (Some(_), Some(_)) => {
248
+ return Err(PyValueError::new_err(
249
+ "Client needs exactly one of `token` or `token_provider`, not both",
250
+ ));
251
+ }
252
+ (Some(t), None) => DbClient::new(&host, &warehouse_id, &t),
238
253
  (None, Some(callable)) => {
239
254
  let provider: Arc<dyn TokenProvider> = Arc::new(PyTokenProvider {
240
255
  callable,
@@ -257,32 +272,6 @@ impl PyDbClient {
257
272
  })
258
273
  }
259
274
 
260
- /// Full submit->poll->fetch->reorder->decode pipeline for one
261
- /// ARROW_STREAM statement, eagerly assembling the whole result. Returns
262
- /// a `pyo3_arrow` `Table` -- it implements the Arrow C Data Interface
263
- /// (`__arrow_c_stream__`), so DuckDB/pyarrow/arro3 can all import it
264
- /// directly, zero-copy. For a large result where you don't want the
265
- /// whole thing pulled upfront, use `execute()` + `ResultSet.fetchmany_arrow`.
266
- #[pyo3(signature = (statement, catalog=None, schema=None, parameters=None))]
267
- fn execute_arrow<'py>(
268
- &self,
269
- py: Python<'py>,
270
- statement: String,
271
- catalog: Option<String>,
272
- schema: Option<String>,
273
- parameters: Option<Py<PyAny>>,
274
- ) -> PyResult<Bound<'py, PyAny>> {
275
- let client = self.inner.clone();
276
- let parameters = parameters_to_value(py, parameters)?;
277
- pyo3_async_runtimes::tokio::future_into_py(py, async move {
278
- let result = pipeline::run_pipeline(client, &statement, catalog.as_deref(), schema.as_deref(), parameters)
279
- .await
280
- .map_err(|e| PyRuntimeError::new_err(e.message))?;
281
- let arrow_schema = result.schema.unwrap_or_else(|| Arc::new(Schema::empty()));
282
- PyTable::try_new(result.batches, arrow_schema).map_err(|e| PyRuntimeError::new_err(e.to_string()))
283
- })
284
- }
285
-
286
275
  /// Submits the statement and starts background chunk fetching, without
287
276
  /// pulling any of it yet -- returns a `ResultSet` for on-demand
288
277
  /// `fetchmany_arrow`/`fetchall_arrow`, mirroring `cursor.py`'s
@@ -312,60 +301,6 @@ impl PyDbClient {
312
301
  })
313
302
  }
314
303
 
315
- /// Like `execute()`, but yields `HEARTBEAT` while waiting on Databricks
316
- /// instead of blocking silently -- for bridging e.g. an SSE connection
317
- /// during a possible multi-minute cold warehouse start. Yields
318
- /// `HEARTBEAT` zero or more times, then a `ResultSet` once ready to
319
- /// fetch. Not async itself (matches `cursor.py`'s own `execute_streamed`,
320
- /// a sync method returning an async generator) -- the submit/poll starts
321
- /// running immediately in the background regardless of when the caller
322
- /// starts iterating.
323
- #[pyo3(signature = (statement, catalog=None, schema=None, parameters=None, total_timeout_s=None))]
324
- fn execute_streamed(
325
- &self,
326
- py: Python<'_>,
327
- statement: String,
328
- catalog: Option<String>,
329
- schema: Option<String>,
330
- parameters: Option<Py<PyAny>>,
331
- total_timeout_s: Option<f64>,
332
- ) -> PyResult<PyExecuteStreamedIter> {
333
- let client = self.inner.clone();
334
- let parameters = parameters_to_value(py, parameters)?;
335
- let fut = async move {
336
- pipeline::execute_lazy(client, &statement, catalog.as_deref(), schema.as_deref(), parameters).await
337
- };
338
- Ok(PyExecuteStreamedIter {
339
- wait: Arc::new(AsyncMutex::new(Some(HeartbeatWait::new(fut, total_timeout_s)))),
340
- })
341
- }
342
-
343
- /// Full submit->poll->fetch->reorder pipeline for one JSON_ARRAY
344
- /// statement -- no Arrow parse at all. Returns a plain list of rows,
345
- /// each row a list of values where every non-null value is a *string*
346
- /// (Databricks' own JSON_ARRAY contract, not this crate's choice) --
347
- /// cast by the manifest's column type_name yourself if you want native
348
- /// Python types.
349
- #[pyo3(signature = (statement, catalog=None, schema=None, parameters=None))]
350
- fn execute_json<'py>(
351
- &self,
352
- py: Python<'py>,
353
- statement: String,
354
- catalog: Option<String>,
355
- schema: Option<String>,
356
- parameters: Option<Py<PyAny>>,
357
- ) -> PyResult<Bound<'py, PyAny>> {
358
- let client = self.inner.clone();
359
- let parameters = parameters_to_value(py, parameters)?;
360
- pyo3_async_runtimes::tokio::future_into_py(py, async move {
361
- let result =
362
- pipeline::run_json_pipeline(client, &statement, catalog.as_deref(), schema.as_deref(), parameters)
363
- .await
364
- .map_err(|e| PyRuntimeError::new_err(e.message))?;
365
- Ok(result.rows)
366
- })
367
- }
368
-
369
304
  /// Uploads `data` to a Unity Catalog volume path via the Files API,
370
305
  /// overwriting anything already there.
371
306
  #[pyo3(signature = (volume_path, data))]
@@ -398,15 +333,16 @@ impl PyDbClient {
398
333
  })
399
334
  }
400
335
 
401
- /// Chunk-at-a-time counterpart to `execute_arrow`, entirely backing
402
- /// `stream_query_json`: yields `HEARTBEAT` while waiting on each
336
+ /// Chunk-at-a-time counterpart to `execute`+`fetchall_arrow`, entirely
337
+ /// backing `stream_query_json`: yields `HEARTBEAT` while waiting on each
403
338
  /// still-in-flight chunk (not just the initial statement wait), then a
404
339
  /// `list[str]` of NDJSON lines (one per row, arro3-`write_ndjson(
405
340
  /// explicit_nulls=True)`-compatible) per chunk as it arrives in logical
406
341
  /// order -- decode and JSON-encoding both happen here, so there's no
407
342
  /// further Python-side conversion step. Not async itself, same as
408
- /// `execute_streamed` -- the returned iterator's `__anext__` does the
409
- /// real work, including the initial submit/poll on its very first call.
343
+ /// `ResultSet.fetchall_arrow_streamed` -- the returned iterator's
344
+ /// `__anext__` does the real work, including the initial submit/poll on
345
+ /// its very first call.
410
346
  #[pyo3(signature = (statement, catalog=None, schema=None, parameters=None, total_timeout_s=None, non_finite_as_string=false))]
411
347
  #[allow(clippy::too_many_arguments)]
412
348
  fn stream_ndjson_lines(
@@ -525,49 +461,6 @@ impl PyResultSet {
525
461
  }
526
462
  }
527
463
 
528
- /// Async iterator returned by `Client.execute_streamed`: yields the
529
- /// `HEARTBEAT` singleton zero or more times, then a `ResultSet` exactly
530
- /// once, then stops.
531
- #[pyclass(name = "ExecuteStreamedIter")]
532
- struct PyExecuteStreamedIter {
533
- wait: Arc<AsyncMutex<Option<HeartbeatWait<ResultStream>>>>,
534
- }
535
-
536
- #[pymethods]
537
- impl PyExecuteStreamedIter {
538
- fn __aiter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> {
539
- slf
540
- }
541
-
542
- fn __anext__<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
543
- let wait = self.wait.clone();
544
- pyo3_async_runtimes::tokio::future_into_py(py, async move {
545
- let mut guard = wait.lock().await;
546
- let Some(w) = guard.as_mut() else {
547
- return Err(PyStopAsyncIteration::new_err(()));
548
- };
549
- match w.tick().await {
550
- Ok(Some(Tick::Heartbeat)) => Python::attach(|py| heartbeat_singleton(py).map(|h| h.into_any())),
551
- Ok(Some(Tick::Ready(stream))) => {
552
- *guard = None;
553
- let result_set = PyResultSet {
554
- statement_id: stream.statement_id.clone(),
555
- num_chunks: stream.num_chunks,
556
- columns: column_pairs(&stream.columns),
557
- inner: Arc::new(AsyncMutex::new(stream)),
558
- };
559
- Python::attach(|py| Py::new(py, result_set).map(|rs| rs.into_any()))
560
- }
561
- Ok(None) => Err(PyStopAsyncIteration::new_err(())),
562
- Err(e) => {
563
- *guard = None;
564
- Err(PyRuntimeError::new_err(e.message))
565
- }
566
- }
567
- })
568
- }
569
- }
570
-
571
464
  /// Async iterator returned by `ResultSet.fetchall_arrow_streamed`: yields
572
465
  /// the `HEARTBEAT` singleton zero or more times, then a `Table` exactly
573
466
  /// once, then stops.
@@ -717,13 +610,11 @@ impl PyNdjsonStreamIter {
717
610
  /// "arrowbricks._core", manifest-path pointing back at this crate).
718
611
  #[pymodule]
719
612
  fn _core(m: &Bound<'_, PyModule>) -> PyResult<()> {
720
- m.add_function(wrap_pyfunction!(ping, m)?)?;
721
613
  m.add_function(wrap_pyfunction!(write_ipc_stream, m)?)?;
722
614
  m.add_function(wrap_pyfunction!(read_ipc_stream, m)?)?;
723
615
  m.add_class::<PyDbClient>()?;
724
616
  m.add_class::<PyResultSet>()?;
725
617
  m.add_class::<PyHeartbeat>()?;
726
- m.add_class::<PyExecuteStreamedIter>()?;
727
618
  m.add_class::<PyFetchallArrowStreamedIter>()?;
728
619
  m.add_class::<PyNdjsonStreamIter>()?;
729
620
  m.add("HEARTBEAT", heartbeat_singleton(m.py())?)?;