From 858704d837053633f77426b89453de0ee0ee8ade Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 10:54:14 -0700 Subject: [PATCH 01/31] Add an ADBC engine adapter xql.register now accepts any ADBC DBAPI connection, so PostgreSQL, SQLite, Snowflake, BigQuery, Flight SQL, DuckDB, and every other database with an ADBC driver share one code path. An ADBC database cannot call back into Python mid-query, so registration streams the Dataset into a table via bulk ingest, reusing the prefetching XarrayPushdownDataset scan for bounded memory. Results round-trip through the existing to_dataset path, which already accepts an ADBC cursor. - mode= ("create" by default, so an existing table is never dropped silently) and temporary= pass through to adbc_ingest. - Mixed-dimension Datasets go into a schema named after the Dataset, keeping name.group portable; databases without schemas (SQLite) and temporary tables fall back to flat name_group tables. - New `adbc` extra; tests run against the SQLite driver and DuckDB's built-in ADBC driver. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- README.md | 7 +- docs/engines.md | 85 ++++++++++-- docs/limitations.md | 21 +++ pyproject.toml | 7 +- tests/test_adbc_backend.py | 225 ++++++++++++++++++++++++++++++++ uv.lock | 196 +++++++++++++++++++--------- xarray_sql/backends/__init__.py | 1 + xarray_sql/backends/adbc.py | 200 ++++++++++++++++++++++++++++ xarray_sql/backends/base.py | 9 +- 9 files changed, 674 insertions(+), 77 deletions(-) create mode 100644 tests/test_adbc_backend.py create mode 100644 xarray_sql/backends/adbc.py diff --git a/README.md b/README.md index 8bca8865..2c322950 100644 --- a/README.md +++ b/README.md @@ -77,10 +77,15 @@ rel = con.sql('SELECT time, AVG("air") AS air FROM air GROUP BY time ORDER BY ti xql.to_dataset(rel, template=ds) # any engine's Arrow result round-trips ``` +Any database with an [ADBC](https://arrow.apache.org/adbc/) driver +(PostgreSQL, SQLite, Snowflake, BigQuery, ...) works too: `xql.register` +ingests the Dataset into a table there, and `xql.to_dataset(cursor, ...)` +brings results back. + `table_names` (below) works the same way on every engine, so a query written against `era5.surface` is not tied to the engine it was written for. -See [Engines](https://xqlsystems.github.io/xarray-sql/latest/engines/) for the support matrix, DuckDB/Polars details, +See [Engines](https://xqlsystems.github.io/xarray-sql/latest/engines/) for the support matrix, DuckDB/Polars/ADBC details, and the lazy chunked round-trip. ## A bigger example: ARCO-ERA5 diff --git a/docs/engines.md b/docs/engines.md index 15a78ae7..d77fc344 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -190,22 +190,85 @@ chunked round-trip is fully supported: windows re-execute on Polars' streaming engine. +## ADBC (adapter; any database with a driver) + +```sh +pip install 'xarray-sql[adbc]' adbc-driver-postgresql # or -sqlite, -snowflake, ... +``` + +[ADBC](https://arrow.apache.org/adbc/) is a database-neutral API whose +drivers speak Arrow natively. `xql.register` accepts any ADBC DBAPI +connection, so PostgreSQL, SQLite, Snowflake, BigQuery, Flight SQL, +DuckDB, and every other database with an ADBC driver share one code +path: + +```python +import adbc_driver_postgresql.dbapi +import xarray_sql as xql + +con = adbc_driver_postgresql.dbapi.connect("postgresql://localhost/weather") +xql.register(con, "era5", ds) # seam 1: ingest +con.commit() + +cur = con.cursor() +cur.execute(""" + SELECT time, lat, lon, AVG(t2m) AS t2m + FROM era5 + WHERE lat BETWEEN 40 AND 41 + GROUP BY time, lat, lon + ORDER BY time, lat, lon +""") +out = xql.to_dataset(cur, template=ds) # seam 2 +``` + +**Registration copies the data.** An ADBC database usually runs in +another process or on another machine, so it cannot call back into +Python to scan a lazy Dataset while a query runs. The adapter instead +streams the Dataset into a new table with ADBC's bulk ingest: chunks +are read on the same prefetching scan the DuckDB adapter uses +(`batch_size`, `prefetch`, `prefetch_bytes`, `coalesce_rows` tune it), +so memory stays bounded while the driver writes, and queries afterwards +run entirely in the database. Ingest what you intend to query — +`ds.sel(...)` a region or `ds[[...]]` a few variables first — rather +than a whole archive. + +Options specific to this adapter: + +- `mode="create"` (default) raises if the table exists; `"replace"` + drops and recreates it; `"append"` and `"create_append"` add rows, + which is how to load a long time series in slices. +- `temporary=True` creates temporary tables that the database drops + when the connection closes — the closest match to the other + engines' register-for-this-session behavior. +- Ingest runs inside the connection's current transaction. The tables + are visible to this connection at once; call `con.commit()` for + other connections to see them. + +Mixed-dimension Datasets are ingested into a database schema named +after the Dataset, so `era5.surface` is the same SQL here as on +DataFusion and DuckDB. On databases without schemas (SQLite), and for +temporary tables, the groups are created as flat `era5_surface` tables +instead (with a warning in the first case). + +The cursor is a one-shot Arrow stream: `xql.to_dataset(cur, ...)` +round-trips eagerly, and `chunks=` needs `spill=True`. + ## Engine support matrix What each integration provides. Known issues and constraints live on [Known issues & limitations](limitations.md). -| | DataFusion | DuckDB | Polars | -|---|---|---|---| -| Register | `XarrayContext` / any `SessionContext` | `xql.register(con, name, ds)` | `pl.scan_pyarrow_dataset(xql.arrow_dataset(ds))` | -| Projection pushdown | yes | yes | yes | -| Chunk pruning on dim predicates | yes | yes | yes | -| Eager round-trip (`xql.to_dataset`) | yes | yes | yes | -| Chunked round-trip (`chunks=`) | re-execution | `spill=True` [^spill-only] | re-execution (streaming engine) | -| `geometry` column ([geospatial](geospatial.md#geoarrow-point-geometry-columns)) | annotated WKB passes through | native `GEOMETRY` (`"wkb"` encoding) | plain binary/struct | -| Mixed-dimension datasets | one schema, `name.group` tables | `name.group` views over `name_group` tables | `xql.arrow_datasets(ds, name)`, one per group | -| Naming those tables (`table_names=`) | yes | yes | yes | -| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | +| | DataFusion | DuckDB | Polars | ADBC | +|---|---|---|---|---| +| Register | `XarrayContext` / any `SessionContext` | `xql.register(con, name, ds)` | `pl.scan_pyarrow_dataset(xql.arrow_dataset(ds))` | `xql.register(con, name, ds)` (copies into the database) | +| Projection pushdown | yes | yes | yes | n/a (the database's own tables) | +| Chunk pruning on dim predicates | yes | yes | yes | n/a (the database's own indexes) | +| Eager round-trip (`xql.to_dataset`) | yes | yes | yes | yes (pass the cursor) | +| Chunked round-trip (`chunks=`) | re-execution | `spill=True` [^spill-only] | re-execution (streaming engine) | `spill=True` | +| `geometry` column ([geospatial](geospatial.md#geoarrow-point-geometry-columns)) | annotated WKB passes through | native `GEOMETRY` (`"wkb"` encoding) | plain binary/struct | driver-dependent | +| Mixed-dimension datasets | one schema, `name.group` tables | `name.group` views over `name_group` tables | `xql.arrow_datasets(ds, name)`, one per group | `name.group` tables in a schema; `name_group` without schemas | +| Naming those tables (`table_names=`) | yes | yes | yes | yes | +| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.0` (tested on 1.12 with SQLite and DuckDB drivers) | [^spill-only]: Why DuckDB relations do not re-execute — and two other engine-specific issues worth knowing — is explained on diff --git a/docs/limitations.md b/docs/limitations.md index d7bc0d18..a043a118 100644 --- a/docs/limitations.md +++ b/docs/limitations.md @@ -82,6 +82,27 @@ Pick your engine: not from the consumer. Not a bug — worth knowing when sizing scans. +=== "ADBC" + + **Registration is a copy, not a lazy view.** + + - *Symptom:* registering a large Dataset takes as long as reading + all of it, and the database holds a full copy. + - *Why:* an ADBC database cannot call back into Python while a + query runs, so there is no lazy scan to push predicates into. + - *What to do:* select the region and variables you will query + before registering; use `mode="append"` to load in slices. + + **SQLite has no timestamp type.** + + - *Symptom:* a datetime dimension comes back from SQLite as text. + The eager round-trip recovers it from the template, but the + chunked round-trip (`chunks=..., spill=True`) cannot build its + window predicates against a text column and fails. + - *What to do:* use the eager round-trip with SQLite, or a database + with native timestamps (PostgreSQL, DuckDB, Snowflake, ...) for + the chunked one. SQLite also widens `float32` to `float64`. + ## Constraints in any engine These follow from the data model — no engine or configuration avoids diff --git a/pyproject.toml b/pyproject.toml index c47f2dfa..b11f02b5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,6 +36,10 @@ dependencies = [ ] [project.optional-dependencies] +adbc = [ + # Plus a driver package for your database, e.g. adbc-driver-postgresql. + "adbc-driver-manager>=1.0", +] duckdb = [ "duckdb>=1.4.0", ] @@ -49,8 +53,9 @@ geo = [ "pyproj", ] test = [ + "adbc-driver-sqlite>=1.0", "cftime", - "xarray-sql[duckdb,polars,geo]", + "xarray-sql[adbc,duckdb,polars,geo]", "pytest", "xarray[io]", "gcsfs", diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py new file mode 100644 index 00000000..3ee7497d --- /dev/null +++ b/tests/test_adbc_backend.py @@ -0,0 +1,225 @@ +"""Tests for the ADBC engine adapter. + +``xql.register`` ingests a Dataset into any database reachable through an +ADBC driver, and ``xql.to_dataset`` rebuilds a labeled Dataset from the +driver's Arrow cursor. Two drivers cover the two shapes of database: +SQLite, which has no schemas and stores timestamps as text, and DuckDB's +built-in ADBC driver, which has schemas and native temporal types. +""" + +import importlib.util + +import numpy as np +import pandas as pd +import pytest +import xarray as xr + +import xarray_sql as xql + +dbapi = pytest.importorskip("adbc_driver_manager.dbapi") +sqlite_dbapi = pytest.importorskip("adbc_driver_sqlite.dbapi") + +NAMES = { + ("time", "lat", "lon"): "surface", + ("time", "level", "lat", "lon"): "atmosphere", +} + + +def _duckdb_driver_path() -> str | None: + """Path of the shared library holding DuckDB's ADBC entrypoint.""" + for module in ("_duckdb", "duckdb.duckdb"): + try: + spec = importlib.util.find_spec(module) + except ModuleNotFoundError: + continue + if spec is not None and spec.origin: + return spec.origin + return None + + +def _connect(driver: str): + if driver == "sqlite": + return sqlite_dbapi.connect() + path = _duckdb_driver_path() + if path is None: + pytest.skip("duckdb is not installed") + return dbapi.connect(driver=path, entrypoint="duckdb_adbc_init") + + +@pytest.fixture(params=["sqlite", "duckdb"]) +def con(request): + connection = _connect(request.param) + yield connection + connection.close() + + +@pytest.fixture +def sqlite_con(): + connection = _connect("sqlite") + yield connection + connection.close() + + +@pytest.fixture +def duckdb_con(): + connection = _connect("duckdb") + yield connection + connection.close() + + +@pytest.fixture +def ds() -> xr.Dataset: + np.random.seed(3) + return xr.Dataset( + data_vars=dict( + temperature=(["time", "lat", "lon"], np.random.randn(8, 5, 6)), + precipitation=(["time", "lat", "lon"], np.random.rand(8, 5, 6)), + ), + coords=dict( + time=pd.date_range("2021-01-01", periods=8, freq="h"), + lat=np.linspace(-10.0, 10.0, 5), + lon=np.linspace(0.0, 40.0, 6), + ), + attrs=dict(description="Synthetic weather."), + ).chunk({"time": 4}) + + +@pytest.fixture +def mixed_ds() -> xr.Dataset: + np.random.seed(11) + return xr.Dataset( + { + "t2m": (["time", "lat", "lon"], np.random.rand(6, 3, 4)), + "temperature": ( + ["time", "level", "lat", "lon"], + np.random.rand(6, 2, 3, 4), + ), + }, + coords={ + "time": pd.date_range("2020-01-01", periods=6, freq="D"), + "lat": np.linspace(-90, 90, 3), + "lon": np.linspace(-180, 180, 4), + "level": [500, 1000], + }, + ).chunk({"time": 2}) + + +def _query(con, sql: str): + cur = con.cursor() + cur.execute(sql) + return cur + + +def test_full_scan_round_trips(con, ds): + xql.register(con, "weather", ds) + + cur = _query( + con, + "SELECT time, lat, lon, temperature, precipitation FROM weather " + "ORDER BY time, lat, lon", + ) + out = xql.to_dataset(cur, template=ds) + + xr.testing.assert_allclose(out, ds.compute()) + assert out.attrs == ds.attrs + + +def test_aggregation_round_trips_on_surviving_dims(con, ds): + xql.register(con, "weather", ds) + + cur = _query( + con, + "SELECT lat, lon, AVG(temperature) AS temperature FROM weather " + "GROUP BY lat, lon ORDER BY lat, lon", + ) + out = xql.to_dataset(cur, template=ds) + + assert out.temperature.dims == ("lat", "lon") + xr.testing.assert_allclose( + out.temperature, ds.temperature.mean("time").compute() + ) + + +def test_chunked_round_trip_spills_the_cursor(duckdb_con, ds): + # DuckDB keeps `time` a timestamp; SQLite returns it as text, which + # only the eager round-trip recovers. + xql.register(duckdb_con, "weather", ds) + + cur = _query( + duckdb_con, + "SELECT time, lat, lon, temperature FROM weather " + "ORDER BY time, lat, lon", + ) + out = xql.to_dataset(cur, template=ds, chunks={"time": 2}, spill=True) + + assert out.temperature.chunks is not None + xr.testing.assert_allclose( + out.temperature.compute(), ds.temperature.compute() + ) + + +def test_existing_table_is_not_overwritten_by_default(con, ds): + xql.register(con, "weather", ds) + + with pytest.raises(dbapi.Error): + xql.register(con, "weather", ds) + + +def test_replace_mode_recreates_the_table(con, ds): + xql.register(con, "weather", ds) + xql.register(con, "weather", ds.isel(time=slice(0, 4)), mode="replace") + + count = _query(con, "SELECT COUNT(*) FROM weather").fetchone()[0] + assert count == 4 * 5 * 6 + + +def test_append_mode_adds_rows(con, ds): + xql.register(con, "weather", ds.isel(time=slice(0, 4))) + xql.register(con, "weather", ds.isel(time=slice(4, 8)), mode="append") + + cur = _query( + con, + "SELECT time, lat, lon, temperature, precipitation FROM weather " + "ORDER BY time, lat, lon", + ) + xr.testing.assert_allclose(xql.to_dataset(cur, template=ds), ds.compute()) + + +def test_temporary_table_is_queryable(con, ds): + xql.register(con, "weather", ds, temporary=True) + + count = _query(con, "SELECT COUNT(*) FROM weather").fetchone()[0] + assert count == 8 * 5 * 6 + + +def test_mixed_dimensions_register_in_a_schema(duckdb_con, mixed_ds): + xql.register(duckdb_con, "era5", mixed_ds, table_names=NAMES) + + cur = _query(duckdb_con, "SELECT AVG(t2m) FROM era5.surface") + assert cur.fetchone()[0] == pytest.approx(float(mixed_ds.t2m.mean())) + cur = _query( + duckdb_con, + "SELECT time, level, lat, lon, temperature FROM era5.atmosphere " + "ORDER BY time, level, lat, lon", + ) + out = xql.to_dataset(cur, template=mixed_ds[["temperature"]]) + xr.testing.assert_allclose(out.temperature, mixed_ds.temperature.compute()) + + +def test_mixed_dimensions_fall_back_to_flat_names(sqlite_con, mixed_ds): + with pytest.warns(RuntimeWarning, match="flat"): + xql.register(sqlite_con, "era5", mixed_ds, table_names=NAMES) + + cur = _query(sqlite_con, "SELECT AVG(t2m) FROM era5_surface") + assert cur.fetchone()[0] == pytest.approx(float(mixed_ds.t2m.mean())) + count = _query(sqlite_con, "SELECT COUNT(*) FROM era5_atmosphere") + assert count.fetchone()[0] == 6 * 2 * 3 * 4 + + +def test_temporary_mixed_dimensions_use_flat_names(duckdb_con, mixed_ds): + xql.register( + duckdb_con, "era5", mixed_ds, table_names=NAMES, temporary=True + ) + + count = _query(duckdb_con, "SELECT COUNT(*) FROM era5_surface") + assert count.fetchone()[0] == 6 * 3 * 4 diff --git a/uv.lock b/uv.lock index cf72eb88..97e1aa39 100644 --- a/uv.lock +++ b/uv.lock @@ -9,6 +9,64 @@ resolution-markers = [ "python_full_version < '3.11'", ] +[[package]] +name = "adbc-driver-manager" +version = "1.12.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/9c/f8/ed6475b49a7cf35ea888d5c95e7d4bc9dc6568f9d741f14c0573d622cc1e/adbc_driver_manager-1.12.0.tar.gz", hash = "sha256:45991f0c2de369d330c6a211ca2edbcce6389c5dc81cde70461bdeb6f8f7b268", size = 217579, upload-time = "2026-07-28T00:43:03.512Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9a/53/2c47920ca9a5bf29893294db2ac765e26823eb3246d0071374d29abdc276/adbc_driver_manager-1.12.0-cp310-cp310-macosx_10_15_x86_64.whl", hash = "sha256:ca18599e19a40da990bffe964475ee27523a87bb770a1ffa77f15c6e73790822", size = 599962, upload-time = "2026-07-28T00:41:45.02Z" }, + { url = "https://files.pythonhosted.org/packages/53/8b/b66dec201f2dcb36d1a794afd5f18310c1252cdf6ee84dd9b58e17a526e0/adbc_driver_manager-1.12.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:6166c5a8ea0904d2ab811f575747ade35ce4cabc1c5acc3cc6468ca158d620e9", size = 610987, upload-time = "2026-07-28T00:41:46.946Z" }, + { url = "https://files.pythonhosted.org/packages/01/9e/3617960d056bdc9f2f2cef0ff902b6e3dd767f3a3f232856edcf113a8eac/adbc_driver_manager-1.12.0-cp310-cp310-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:41dadba88e1806eba6cb3eb30b7a2e9f804001bb002dd18ed6a15edb6f5d096f", size = 4596654, upload-time = "2026-07-28T00:41:49.282Z" }, + { url = "https://files.pythonhosted.org/packages/53/a0/224464451cf28baea8033cae16fd1819d5d769ecd72346363de6c3189e3a/adbc_driver_manager-1.12.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:63048664b31c964ae9cc0c1bf3902ec7c26751bee110ab320d78f8d1af7e0b6a", size = 4679094, upload-time = "2026-07-28T00:41:51.455Z" }, + { url = "https://files.pythonhosted.org/packages/35/cf/8089661f92a3991edcd8938c2fe96cbb7a8d1298623aaceafdf78f8ff8cc/adbc_driver_manager-1.12.0-cp310-cp310-win_amd64.whl", hash = "sha256:bf7764d4f1ac9b54e442d6c3b6afbefce639268a7e505a05629507209fe0e3f7", size = 763065, upload-time = "2026-07-28T00:41:53.11Z" }, + { url = "https://files.pythonhosted.org/packages/73/2d/e41ea911f9486c497534ae181dfdab19adca21f71abc8a1fcaf2c27251a8/adbc_driver_manager-1.12.0-cp311-cp311-macosx_10_15_x86_64.whl", hash = "sha256:3c0c73670c8aa6fe42de1d5e71a0b329c4b37f7c55c560c23f6f3a1609200c1f", size = 600030, upload-time = "2026-07-28T00:41:54.753Z" }, + { url = "https://files.pythonhosted.org/packages/10/ea/1a8b51999785d7dce17079dd635c9ee2372c75ec7328423ce862741cb503/adbc_driver_manager-1.12.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:6943c7adcf3c7c9f7c4b5bdb7589c331027a347e3c77471eb3f656b1a881e351", size = 611058, upload-time = "2026-07-28T00:41:56.306Z" }, + { url = "https://files.pythonhosted.org/packages/85/a2/5ede53173a420742fa71d6c26792e2295fc73f25cf39055f912a5385b1f0/adbc_driver_manager-1.12.0-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:78c9936adb280e2c10e90632e41b58aa23be358e1136d8fb3c52862b72818a95", size = 4663735, upload-time = "2026-07-28T00:41:58.793Z" }, + { url = "https://files.pythonhosted.org/packages/1c/1a/781561d0f55e05a0b884244ed563dab14b165b4dd74abd8af3f8efd95e3d/adbc_driver_manager-1.12.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:30d96ab4a2594b4109496fb4913646f41a5bf1ecce79b4313847d240a2a62db3", size = 4747090, upload-time = "2026-07-28T00:42:00.746Z" }, + { url = "https://files.pythonhosted.org/packages/c6/fa/47c755a74ea4887968c52a968e02736007da4042fc0820292dc8c6827a94/adbc_driver_manager-1.12.0-cp311-cp311-win_amd64.whl", hash = "sha256:67419b92c286646944426992069f56fed90c2ceac83521f6d66d7d3cbf6c17ea", size = 760952, upload-time = "2026-07-28T00:42:02.432Z" }, + { url = "https://files.pythonhosted.org/packages/de/8c/cd3fe16df716719116a6c79e64a768fe994f6ded55d5a8f091bb4f42d6f0/adbc_driver_manager-1.12.0-cp312-cp312-macosx_10_15_x86_64.whl", hash = "sha256:fd02364c65b8b376c5627e3b77410f457fcbbf983e52e8d15ca099da3a7ae314", size = 599054, upload-time = "2026-07-28T00:42:04.072Z" }, + { url = "https://files.pythonhosted.org/packages/49/4a/2f060ff6bd61420ea1613670e1f85a22a8714934c235186dc3803de8ddac/adbc_driver_manager-1.12.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:d8dcf62621090e8d9c8216e08dfc4043f16331872522186af61a5de9478e9c63", size = 609964, upload-time = "2026-07-28T00:42:05.82Z" }, + { url = "https://files.pythonhosted.org/packages/8a/f1/0746db149828ae91e4a6cf49f8d0e49210eec20c03ad80044454139c8240/adbc_driver_manager-1.12.0-cp312-cp312-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:efa5dbbf101962d212b176f25e6fc509dacf07afd4cf70b5027d81ec6871bdec", size = 4685726, upload-time = "2026-07-28T00:42:08.1Z" }, + { url = "https://files.pythonhosted.org/packages/b9/c3/f8e9c5157b19e986df719259eb3502dad1268df9f7a1034f65ca220ab2ea/adbc_driver_manager-1.12.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:8b340679a005a8adf6b0b58754dbc638dff00db7b2559c140406a1d92678b48c", size = 4768774, upload-time = "2026-07-28T00:42:10.359Z" }, + { url = "https://files.pythonhosted.org/packages/92/51/f8e625af691e6b4c54945790854524356a02a0a69063e888f7cfee1b2e50/adbc_driver_manager-1.12.0-cp312-cp312-win_amd64.whl", hash = "sha256:47f428a922d224fd486b661deeaf9520e5faec558b3d144832bed09a080cac88", size = 760087, upload-time = "2026-07-28T00:42:11.871Z" }, + { url = "https://files.pythonhosted.org/packages/9a/f9/674c5bbc5093617d72c4f58a5dab67982710b2320cc9aa826050a6aaa131/adbc_driver_manager-1.12.0-cp313-cp313-macosx_10_15_x86_64.whl", hash = "sha256:c42ca4d9caa22b3a5ce76bde8729169f403bb7393e3671734b9416634c207125", size = 596815, upload-time = "2026-07-28T00:42:13.64Z" }, + { url = "https://files.pythonhosted.org/packages/56/5f/c1d888d787330801edae282d2a9def3765e8157547cc20e71154ff38c1bb/adbc_driver_manager-1.12.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:c894117c8f5c484b902c8b070bcfd9d31d90efe0288b2b58a3ddab97c80f66e7", size = 608277, upload-time = "2026-07-28T00:42:15.643Z" }, + { url = "https://files.pythonhosted.org/packages/06/4b/ee799babf171e39690ef45560451096f869d9e7387bc0e5a754bb243ed2a/adbc_driver_manager-1.12.0-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:214f80f9b65562f08b4d1c52a756b5db557530e3c0652f587c43aaa80039579a", size = 4667230, upload-time = "2026-07-28T00:42:17.97Z" }, + { url = "https://files.pythonhosted.org/packages/00/c6/a35e38ef5e0db391be79e0e14c019ce378b87d9d7e31d1dfcd451e9d291f/adbc_driver_manager-1.12.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:532ab290b3d923ce0a75bca21dc6e13f55835625f78808e1664755939f3ebdf6", size = 4745299, upload-time = "2026-07-28T00:42:20.189Z" }, + { url = "https://files.pythonhosted.org/packages/16/e2/62bacd6844859036d79ea229401b5200056fb5050c82dc3a2e28b08ff49b/adbc_driver_manager-1.12.0-cp313-cp313-win_amd64.whl", hash = "sha256:034da82c1a6e195d67ca1f0c97a1a517046037ec3029ab9a0ea8f7ccb14056e4", size = 758878, upload-time = "2026-07-28T00:42:21.598Z" }, + { url = "https://files.pythonhosted.org/packages/50/ea/f53b434fe36d0f138d147fc10a95784c8c0eeea1bec1f3f31eee5ec8bdb5/adbc_driver_manager-1.12.0-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:a740d634118722f42af31176374fddbad3846fa2e6536f497bac145e9511cecc", size = 597579, upload-time = "2026-07-28T00:42:23.216Z" }, + { url = "https://files.pythonhosted.org/packages/ba/57/6208e66d9256550c2aff75db4a323a855a0d5d2d1bd639526f825d3e08b4/adbc_driver_manager-1.12.0-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:8a77ae39832e67946009816d83c321e540a3024aad1419ccba24ddeb7b6a01f4", size = 610337, upload-time = "2026-07-28T00:42:25.051Z" }, + { url = "https://files.pythonhosted.org/packages/1d/cd/f5ea3f08191af5ae15041821fcb52bf35837dce1a9ac16fa039b3bfe308c/adbc_driver_manager-1.12.0-cp314-cp314-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:690f140ca67d49f995afac59f85441c3d5e896cd2fc8fd381423fe900e51f1f7", size = 4664297, upload-time = "2026-07-28T00:42:27.474Z" }, + { url = "https://files.pythonhosted.org/packages/df/81/823a71a515078545eab8a4be8381206887129e11b91e9bf51ca2a9eea44d/adbc_driver_manager-1.12.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:fd568c94874c0586d82f99de2bb5d2c02b4fa9c5bafe3d0d8ab353bddf9d2fd6", size = 4733739, upload-time = "2026-07-28T00:42:29.814Z" }, + { url = "https://files.pythonhosted.org/packages/cf/f7/7612d078d935344aee679a44a6283de6aae9008eb8e0ef80e475dd12dffa/adbc_driver_manager-1.12.0-cp314-cp314-win_amd64.whl", hash = "sha256:57f5101fb2a853b1ffb81ff807b5e29a51ba14c64032eb0038b8dfd433b6d533", size = 777952, upload-time = "2026-07-28T00:42:40.881Z" }, + { url = "https://files.pythonhosted.org/packages/b0/ad/2478338aaece38b8b72259dbfd4d4c84d9a038421e25bbc283e510d47555/adbc_driver_manager-1.12.0-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:bb9db6e4a3bcd73153435a900b5ae40ad36f5875df93a8faf784d9fcf6833983", size = 615694, upload-time = "2026-07-28T00:42:31.932Z" }, + { url = "https://files.pythonhosted.org/packages/bc/a0/0592c85e653f005aa28de7733b3c3c4f0282238301694f76806e5f3cc1e1/adbc_driver_manager-1.12.0-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:07cae26bd5ccee6caa4227f817c0fd57f9ac131c2dd98e0c5d7fecfef61819c7", size = 628341, upload-time = "2026-07-28T00:42:33.481Z" }, + { url = "https://files.pythonhosted.org/packages/9d/00/65705a72f768bc2dda82623a74cf816609dfdff56f3ad22b073d4a1ea7f8/adbc_driver_manager-1.12.0-cp314-cp314t-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:442ed2ee8ea62c475bf3478385555bb4f0b25d9d551087ffe40c73b91bf5431e", size = 4730268, upload-time = "2026-07-28T00:42:35.661Z" }, + { url = "https://files.pythonhosted.org/packages/44/b9/60ecde5d9dde5acc5576cb0ba5ffa34e154464e07fa295c57cd975ea27c7/adbc_driver_manager-1.12.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:9c2aa05c5dc52164692284b2df27fba5680dbc967b8e3ca704aabf5399667996", size = 4777527, upload-time = "2026-07-28T00:42:37.709Z" }, + { url = "https://files.pythonhosted.org/packages/ac/76/6749e0c0c437219780c65487cff67dc09a556c1fccf577a2b27f7b92a704/adbc_driver_manager-1.12.0-cp314-cp314t-win_amd64.whl", hash = "sha256:cfa08f8c7c63e3fa92eb4e26ef4d8a9520cf92a39281cd011821f6f16a963080", size = 793451, upload-time = "2026-07-28T00:42:39.222Z" }, +] + +[[package]] +name = "adbc-driver-sqlite" +version = "1.12.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "adbc-driver-manager" }, + { name = "importlib-resources" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/f5/02/2dc143bdd2a62c52d103d4b0ae491a347944aaed25b3d40fb11797750c70/adbc_driver_sqlite-1.12.0.tar.gz", hash = "sha256:18466a2f0c14f94cb0b17818157cc14ed6b93aef0a48ef648de945e9bac1540d", size = 12849, upload-time = "2026-07-28T00:43:05.408Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a1/f7/c35740269d3a5e3aa07b9ab155d4e943a7f5267f64d8f7396d5e14184a02/adbc_driver_sqlite-1.12.0-py3-none-macosx_10_15_x86_64.whl", hash = "sha256:2d5b3e9d0b5dbc66324b0ccf2ded886e3781f901be986892d319529b05536d3b", size = 1413592, upload-time = "2026-07-28T00:42:53.523Z" }, + { url = "https://files.pythonhosted.org/packages/e6/31/5d1d637e6ae76fcc57d5116d537aa78d2ab687354d5a0b2d527e085b61d4/adbc_driver_sqlite-1.12.0-py3-none-macosx_11_0_arm64.whl", hash = "sha256:5a81f53791e4aec69afbf8f77dac6acf48749fd84684e86601eafdd36d2eb7c3", size = 1357479, upload-time = "2026-07-28T00:42:55.194Z" }, + { url = "https://files.pythonhosted.org/packages/6c/99/415bf90eb912403d2d5d0c31baa1cedf200bd510f40027ee8fd3421c4c02/adbc_driver_sqlite-1.12.0-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:c987d03e3f4850e57f218c8a0b9d224209123af642469ee1f36901c5a51725bd", size = 1501753, upload-time = "2026-07-28T00:42:57.442Z" }, + { url = "https://files.pythonhosted.org/packages/69/10/a3156f19fadd254a4f58a328a8aa9472c981ff93bb4d23f3c22a4341796e/adbc_driver_sqlite-1.12.0-py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:3005a80bedf6624c6856da98037ea943a791aa8e82dad458259e0558be32912c", size = 1548999, upload-time = "2026-07-28T00:42:59.163Z" }, + { url = "https://files.pythonhosted.org/packages/f4/d9/3245d741936100365ea77c434f84a0985523467bd812e5e11bb9b36d7152/adbc_driver_sqlite-1.12.0-py3-none-win_amd64.whl", hash = "sha256:0982bfc06158c2140b5c490b1a1325019c827b158f8432a30d49c8a0c18533ad", size = 1522964, upload-time = "2026-07-28T00:43:00.917Z" }, +] + [[package]] name = "aiohappyeyeballs" version = "2.6.1" @@ -160,8 +218,8 @@ name = "beautifulsoup4" version = "4.13.4" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "soupsieve", marker = "python_full_version >= '3.11'" }, - { name = "typing-extensions", marker = "python_full_version >= '3.11'" }, + { name = "soupsieve" }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/d8/e4/0c4c39e18fd76d6a628d4dd8da40543d136ce2d1752bd6eeeab0791f4d6b/beautifulsoup4-4.13.4.tar.gz", hash = "sha256:dbb3c4e1ceae6aefebdaf2423247260cd062430a410e38c66f2baa50a8437195", size = 621067, upload-time = "2025-04-15T17:05:13.836Z" } wheels = [ @@ -182,8 +240,8 @@ name = "cattrs" version = "25.1.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "attrs", marker = "python_full_version >= '3.11'" }, - { name = "typing-extensions", marker = "python_full_version >= '3.11'" }, + { name = "attrs" }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/57/2b/561d78f488dcc303da4639e02021311728fb7fda8006dd2835550cddd9ed/cattrs-25.1.1.tar.gz", hash = "sha256:c914b734e0f2d59e5b720d145ee010f1fd9a13ee93900922a2f3f9d593b8382c", size = 435016, upload-time = "2025-06-04T20:27:15.44Z" } wheels = [ @@ -477,7 +535,7 @@ name = "donfig" version = "0.8.1.post1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "pyyaml", marker = "python_full_version >= '3.11'" }, + { name = "pyyaml" }, ] sdist = { url = "https://files.pythonhosted.org/packages/25/71/80cc718ff6d7abfbabacb1f57aaa42e9c1552bfdd01e64ddd704e4a03638/donfig-0.8.1.post1.tar.gz", hash = "sha256:3bef3413a4c1c601b585e8d297256d0c1470ea012afa6e8461dc28bfb7c23f52", size = 19506, upload-time = "2024-05-23T14:14:31.513Z" } wheels = [ @@ -531,7 +589,7 @@ name = "exceptiongroup" version = "1.3.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "typing-extensions", marker = "python_full_version < '3.11'" }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/0b/9f/a65090624ecf468cdca03533906e7c69ed7588582240cfe7cc9e770b50eb/exceptiongroup-1.3.0.tar.gz", hash = "sha256:b241f5885f560bc56a59ee63ca4c6a8bfa46ae4ad651af316d4e81817bb9fd88", size = 29749, upload-time = "2025-05-10T17:42:51.123Z" } wheels = [ @@ -898,13 +956,22 @@ name = "importlib-metadata" version = "8.7.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "zipp", marker = "python_full_version < '3.12'" }, + { name = "zipp" }, ] sdist = { url = "https://files.pythonhosted.org/packages/76/66/650a33bd90f786193e4de4b3ad86ea60b53c89b669a5c7be931fac31cdb0/importlib_metadata-8.7.0.tar.gz", hash = "sha256:d13b81ad223b890aa16c5471f2ac3056cf76c5f10f82d6f9292f0b415f389000", size = 56641, upload-time = "2025-04-27T15:29:01.736Z" } wheels = [ { url = "https://files.pythonhosted.org/packages/20/b0/36bd937216ec521246249be3bf9855081de4c5e06a0c9b4219dbeda50373/importlib_metadata-8.7.0-py3-none-any.whl", hash = "sha256:e5dd1551894c77868a30651cef00984d50e1002d06942a7101d34870c5f02afd", size = 27656, upload-time = "2025-04-27T15:29:00.214Z" }, ] +[[package]] +name = "importlib-resources" +version = "7.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/e4/06/b56dfa750b44e86157093bc8fca0ab81dccbf5260510de4eaf1cb69b5b99/importlib_resources-7.1.0.tar.gz", hash = "sha256:0722d4c6212489c530f2a145a34c0a7a3b4721bc96a15fada5930e2a0b760708", size = 44985, upload-time = "2026-04-12T16:36:09.232Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/8a/db/55a262f3606bebcae07cc14095338471ad7c0bbcaa37707e6f0ee49725b7/importlib_resources-7.1.0-py3-none-any.whl", hash = "sha256:1bd7b48b4088eddb2cd16382150bb515af0bd2c70128194392725f82ad2c96a1", size = 37232, upload-time = "2026-04-12T16:36:08.219Z" }, +] + [[package]] name = "iniconfig" version = "2.1.0" @@ -1397,7 +1464,7 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" } }, ] sdist = { url = "https://files.pythonhosted.org/packages/85/56/8895a76abe4ec94ebd01eeb6d74f587bc4cddd46569670e1402852a5da13/numcodecs-0.13.1.tar.gz", hash = "sha256:a3cf37881df0898f3a9c0d4477df88133fe85185bffe57ba31bcc2fa207709bc", size = 5955215, upload-time = "2024-10-09T16:28:00.188Z" } wheels = [ @@ -1430,8 +1497,8 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, - { name = "typing-extensions", marker = "python_full_version >= '3.11'" }, + { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" } }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/00/35/49da850ce5371da3930d099da364a73ce9ae4fc64075e521674b48f4804d/numcodecs-0.16.1.tar.gz", hash = "sha256:c47f20d656454568c6b4697ce02081e6bbb512f198738c6a56fafe8029c97fb1", size = 6268134, upload-time = "2025-05-22T13:33:04.098Z" } wheels = [ @@ -1454,7 +1521,7 @@ wheels = [ [package.optional-dependencies] crc32c = [ - { name = "crc32c", marker = "python_full_version >= '3.11'" }, + { name = "crc32c" }, ] [[package]] @@ -1964,13 +2031,13 @@ name = "pydap" version = "3.5.5" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "beautifulsoup4", marker = "python_full_version >= '3.11'" }, - { name = "lxml", marker = "python_full_version >= '3.11'" }, - { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, - { name = "requests", marker = "python_full_version >= '3.11'" }, - { name = "requests-cache", marker = "python_full_version >= '3.11'" }, - { name = "scipy", version = "1.16.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, - { name = "webob", marker = "python_full_version >= '3.11'" }, + { name = "beautifulsoup4" }, + { name = "lxml" }, + { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" } }, + { name = "requests" }, + { name = "requests-cache" }, + { name = "scipy", version = "1.16.0", source = { registry = "https://pypi.org/simple" } }, + { name = "webob" }, ] sdist = { url = "https://files.pythonhosted.org/packages/36/14/4faa1b6cf6e051e362467992aac43fbbe64aab2b8316528cddb0fec5de5d/pydap-3.5.5.tar.gz", hash = "sha256:0f8ca9b4e244c4d345d0b5269c4ebc886fcd0778b828e5ae1415b7ea5341eabd", size = 13011746, upload-time = "2025-04-14T06:02:00.326Z" } wheels = [ @@ -2007,7 +2074,7 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "certifi", marker = "python_full_version < '3.11'" }, + { name = "certifi" }, ] sdist = { url = "https://files.pythonhosted.org/packages/67/10/a8480ea27ea4bbe896c168808854d00f2a9b49f95c0319ddcbba693c8a90/pyproj-3.7.1.tar.gz", hash = "sha256:60d72facd7b6b79853f19744779abcd3f804c4e0d4fa8815469db20c9f640a47", size = 226339, upload-time = "2025-02-16T04:28:46.621Z" } wheels = [ @@ -2056,7 +2123,7 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "certifi", marker = "python_full_version >= '3.11'" }, + { name = "certifi" }, ] sdist = { url = "https://files.pythonhosted.org/packages/04/90/67bd7260b4ea9b8b20b4f58afef6c223ecb3abf368eb4ec5bc2cdef81b49/pyproj-3.7.2.tar.gz", hash = "sha256:39a0cf1ecc7e282d1d30f36594ebd55c9fae1fda8a2622cee5d100430628f88c", size = 226279, upload-time = "2025-08-14T12:05:42.18Z" } wheels = [ @@ -2244,12 +2311,12 @@ name = "requests-cache" version = "1.2.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "attrs", marker = "python_full_version >= '3.11'" }, - { name = "cattrs", marker = "python_full_version >= '3.11'" }, - { name = "platformdirs", marker = "python_full_version >= '3.11'" }, - { name = "requests", marker = "python_full_version >= '3.11'" }, - { name = "url-normalize", marker = "python_full_version >= '3.11'" }, - { name = "urllib3", marker = "python_full_version >= '3.11'" }, + { name = "attrs" }, + { name = "cattrs" }, + { name = "platformdirs" }, + { name = "requests" }, + { name = "url-normalize" }, + { name = "urllib3" }, ] sdist = { url = "https://files.pythonhosted.org/packages/1a/be/7b2a95a9e7a7c3e774e43d067c51244e61dea8b120ae2deff7089a93fb2b/requests_cache-1.2.1.tar.gz", hash = "sha256:68abc986fdc5b8d0911318fbb5f7c80eebcd4d01bfacc6685ecf8876052511d1", size = 3018209, upload-time = "2024-06-18T17:18:03.774Z" } wheels = [ @@ -2314,7 +2381,7 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" } }, ] sdist = { url = "https://files.pythonhosted.org/packages/0f/37/6964b830433e654ec7485e45a00fc9a27cf868d622838f6b6d9c5ec0d532/scipy-1.15.3.tar.gz", hash = "sha256:eae3cf522bc7df64b42cad3925c876e1b0b6c35c1337c93e12c0f366f55b0eaf", size = 59419214, upload-time = "2025-05-08T16:13:05.955Z" } wheels = [ @@ -2376,7 +2443,7 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, + { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" } }, ] sdist = { url = "https://files.pythonhosted.org/packages/81/18/b06a83f0c5ee8cddbde5e3f3d0bb9b702abfa5136ef6d4620ff67df7eee5/scipy-1.16.0.tar.gz", hash = "sha256:b5ef54021e832869c8cfb03bc3bf20366cbcd426e02a58e8a58d7584dfbb8f62", size = 30581216, upload-time = "2025-06-22T16:27:55.782Z" } wheels = [ @@ -2507,7 +2574,7 @@ name = "url-normalize" version = "2.2.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "idna", marker = "python_full_version >= '3.11'" }, + { name = "idna" }, ] sdist = { url = "https://files.pythonhosted.org/packages/80/31/febb777441e5fcdaacb4522316bf2a527c44551430a4873b052d545e3279/url_normalize-2.2.1.tar.gz", hash = "sha256:74a540a3b6eba1d95bdc610c24f2c0141639f3ba903501e61a52a8730247ff37", size = 18846, upload-time = "2025-04-26T20:37:58.553Z" } wheels = [ @@ -2694,9 +2761,9 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "packaging", marker = "python_full_version < '3.11'" }, - { name = "pandas", marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" } }, + { name = "packaging" }, + { name = "pandas" }, ] sdist = { url = "https://files.pythonhosted.org/packages/19/ec/e50d833518f10b0c24feb184b209bb6856f25b919ba8c1f89678b930b1cd/xarray-2025.6.1.tar.gz", hash = "sha256:a84f3f07544634a130d7dc615ae44175419f4c77957a7255161ed99c69c7c8b0", size = 3003185, upload-time = "2025-06-12T03:04:09.099Z" } wheels = [ @@ -2705,13 +2772,13 @@ wheels = [ [package.optional-dependencies] io = [ - { name = "cftime", marker = "python_full_version < '3.11'" }, - { name = "fsspec", marker = "python_full_version < '3.11'" }, - { name = "h5netcdf", marker = "python_full_version < '3.11'" }, - { name = "netcdf4", marker = "python_full_version < '3.11'" }, - { name = "pooch", marker = "python_full_version < '3.11'" }, - { name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "zarr", version = "2.18.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "cftime" }, + { name = "fsspec" }, + { name = "h5netcdf" }, + { name = "netcdf4" }, + { name = "pooch" }, + { name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" } }, + { name = "zarr", version = "2.18.3", source = { registry = "https://pypi.org/simple" } }, ] [[package]] @@ -2725,9 +2792,9 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, - { name = "packaging", marker = "python_full_version >= '3.11'" }, - { name = "pandas", marker = "python_full_version >= '3.11'" }, + { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" } }, + { name = "packaging" }, + { name = "pandas" }, ] sdist = { url = "https://files.pythonhosted.org/packages/75/b5/d2f1fb2f583ae803b86280dcb7a4f81eef9d4c54f1ecde2307d1b6b1a147/xarray-2025.7.0.tar.gz", hash = "sha256:fd83ac8d638e7caef9d7f0c82bcdf380cede29d2ff84575133b2b95164af78ee", size = 3005754, upload-time = "2025-07-03T16:37:27.427Z" } wheels = [ @@ -2736,14 +2803,14 @@ wheels = [ [package.optional-dependencies] io = [ - { name = "cftime", marker = "python_full_version >= '3.11'" }, - { name = "fsspec", marker = "python_full_version >= '3.11'" }, - { name = "h5netcdf", marker = "python_full_version >= '3.11'" }, - { name = "netcdf4", marker = "python_full_version >= '3.11'" }, - { name = "pooch", marker = "python_full_version >= '3.11'" }, - { name = "pydap", marker = "python_full_version >= '3.11'" }, - { name = "scipy", version = "1.16.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, - { name = "zarr", version = "3.0.10", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, + { name = "cftime" }, + { name = "fsspec" }, + { name = "h5netcdf" }, + { name = "netcdf4" }, + { name = "pooch" }, + { name = "pydap" }, + { name = "scipy", version = "1.16.0", source = { registry = "https://pypi.org/simple" } }, + { name = "zarr", version = "3.0.10", source = { registry = "https://pypi.org/simple" } }, ] [[package]] @@ -2757,6 +2824,9 @@ dependencies = [ ] [package.optional-dependencies] +adbc = [ + { name = "adbc-driver-manager" }, +] dev = [ { name = "mkdocstrings", extra = ["python"] }, { name = "pre-commit" }, @@ -2779,6 +2849,8 @@ polars = [ { name = "polars" }, ] test = [ + { name = "adbc-driver-manager" }, + { name = "adbc-driver-sqlite" }, { name = "cftime" }, { name = "duckdb" }, { name = "gcsfs" }, @@ -2800,6 +2872,8 @@ dev = [ [package.metadata] requires-dist = [ + { name = "adbc-driver-manager", marker = "extra == 'adbc'", specifier = ">=1.0" }, + { name = "adbc-driver-sqlite", marker = "extra == 'test'", specifier = ">=1.0" }, { name = "cftime", marker = "extra == 'test'" }, { name = "dask", specifier = ">=2024.8.0" }, { name = "datafusion", specifier = "==54.0.0" }, @@ -2814,11 +2888,11 @@ requires-dist = [ { name = "watchfiles", marker = "extra == 'dev'" }, { name = "xarray", specifier = ">=2024.7.0" }, { name = "xarray", extras = ["io"], marker = "extra == 'test'" }, + { name = "xarray-sql", extras = ["adbc", "duckdb", "polars", "geo"], marker = "extra == 'test'" }, { name = "xarray-sql", extras = ["docs"], marker = "extra == 'dev'" }, - { name = "xarray-sql", extras = ["duckdb", "polars", "geo"], marker = "extra == 'test'" }, { name = "zensical", marker = "extra == 'docs'" }, ] -provides-extras = ["dev", "docs", "duckdb", "geo", "polars", "test"] +provides-extras = ["adbc", "dev", "docs", "duckdb", "geo", "polars", "test"] [package.metadata.requires-dev] dev = [ @@ -2936,10 +3010,10 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "asciitree", marker = "python_full_version < '3.11'" }, - { name = "fasteners", marker = "python_full_version < '3.11' and sys_platform != 'emscripten'" }, - { name = "numcodecs", version = "0.13.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "asciitree" }, + { name = "fasteners", marker = "sys_platform != 'emscripten'" }, + { name = "numcodecs", version = "0.13.1", source = { registry = "https://pypi.org/simple" } }, + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" } }, ] sdist = { url = "https://files.pythonhosted.org/packages/23/c4/187a21ce7cf7c8f00c060dd0e04c2a81139bb7b1ab178bba83f2e1134ce2/zarr-2.18.3.tar.gz", hash = "sha256:2580d8cb6dd84621771a10d31c4d777dca8a27706a1a89b29f42d2d37e2df5ce", size = 3603224, upload-time = "2024-09-04T23:20:16.595Z" } wheels = [ @@ -2957,11 +3031,11 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "donfig", marker = "python_full_version >= '3.11'" }, - { name = "numcodecs", version = "0.16.1", source = { registry = "https://pypi.org/simple" }, extra = ["crc32c"], marker = "python_full_version >= '3.11'" }, - { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, - { name = "packaging", marker = "python_full_version >= '3.11'" }, - { name = "typing-extensions", marker = "python_full_version >= '3.11'" }, + { name = "donfig" }, + { name = "numcodecs", version = "0.16.1", source = { registry = "https://pypi.org/simple" }, extra = ["crc32c"] }, + { name = "numpy", version = "2.3.1", source = { registry = "https://pypi.org/simple" } }, + { name = "packaging" }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/07/10/a1b6eabeb5a8681916568a7c6a7a1849c952131be127ccbd57e05d47d43e/zarr-3.0.10.tar.gz", hash = "sha256:1fd1318ade646f692d8f604be0e0ad125675a061196e612e3f7a2cfa9e957d1c", size = 263594, upload-time = "2025-07-03T17:29:27.733Z" } wheels = [ diff --git a/xarray_sql/backends/__init__.py b/xarray_sql/backends/__init__.py index af0334be..476056e5 100644 --- a/xarray_sql/backends/__init__.py +++ b/xarray_sql/backends/__init__.py @@ -15,6 +15,7 @@ """ from .base import EngineAdapter, get_adapter, register, register_adapter +from . import adbc as _adbc # noqa: F401 (self-registers) from . import datafusion as _datafusion # noqa: F401 (self-registers) from . import duckdb as _duckdb # noqa: F401 (self-registers) from .pyarrow import ( diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py new file mode 100644 index 00000000..4c826c74 --- /dev/null +++ b/xarray_sql/backends/adbc.py @@ -0,0 +1,200 @@ +"""ADBC engine adapter. + +[ADBC](https://arrow.apache.org/adbc/) (Arrow Database Connectivity) is +a database-neutral API whose drivers speak Arrow natively: PostgreSQL, +SQLite, Snowflake, BigQuery, Flight SQL, DuckDB, and more, each behind +the same DBAPI-style ``Connection``. This adapter registers a Dataset on +any of them through ADBC's bulk-ingest call, so one code path reaches +every database with an ADBC driver. + +Unlike the DataFusion and DuckDB adapters, registration here *copies* +the data: an ADBC database is typically a separate process or service +that cannot call back into Python to scan a lazy Dataset while a query +runs. The Dataset is streamed chunk by chunk (prefetched, with bounded +memory) into a new table, so only the ingest itself touches the source +data, and queries afterwards run entirely inside the database. + +Results come back through ADBC's Arrow cursor API, which +[xarray_sql.to_dataset][] accepts directly:: + + cur = con.cursor() + cur.execute("SELECT time, lat, lon, t2m FROM era5 ORDER BY time") + out = xql.to_dataset(cur, template=ds) + +This adapter never imports ``adbc_driver_manager`` at runtime — +detection is by connection type — so ADBC stays a purely optional +dependency (``pip install xarray-sql[adbc]`` plus a driver package). +""" + +from __future__ import annotations + +import warnings +from typing import TYPE_CHECKING, Any, Literal, TypeGuard + +import xarray as xr + +from ..df import ( + Chunks, + TableNames, + group_vars_by_dims, + resolve_table_names, + shared_coord_arrays, +) +from .base import register_adapter +from .pyarrow import XarrayPushdownDataset + +if TYPE_CHECKING: + from adbc_driver_manager import dbapi + +__all__ = ["ADBCAdapter"] + +IngestMode = Literal["create", "append", "replace", "create_append"] + + +def _quote(identifier: str) -> str: + """Render *identifier* as a quoted SQL identifier.""" + escaped = identifier.replace('"', '""') + return f'"{escaped}"' + + +def _ingest( + con: dbapi.Connection, + table: str, + ds: xr.Dataset, + chunks: Chunks, + *, + mode: IngestMode, + temporary: bool, + db_schema_name: str | None = None, + **kwargs: Any, +) -> None: + """Stream *ds* into *table* through ADBC bulk ingest. + + The full scan of an + [XarrayPushdownDataset][xarray_sql.backends.pyarrow.XarrayPushdownDataset] + prefetches chunks on a thread pool while the driver writes earlier + batches, so the source read and the database write overlap. + """ + reader = XarrayPushdownDataset(ds, chunks, **kwargs).scanner().to_reader() + with con.cursor() as cur: + cur.adbc_ingest( + table, + reader, + mode=mode, + db_schema_name=db_schema_name, + temporary=temporary, + ) + + +def _create_schema(con: dbapi.Connection, name: str) -> bool: + """Create the database schema *name*; whether that succeeded. + + Not every ADBC database has schemas (SQLite does not), and creating + one can need privileges the connection lacks. + """ + try: + with con.cursor() as cur: + cur.execute(f"CREATE SCHEMA IF NOT EXISTS {_quote(name)}") + except Exception as exc: # noqa: BLE001 — fall back to flat names + warnings.warn( + f"Could not create the {name!r} schema to hold the dimension " + f"groups of {name!r} ({exc}); registering them as flat " + f"{name}_ tables instead.", + RuntimeWarning, + stacklevel=4, + ) + return False + return True + + +@register_adapter +class ADBCAdapter: + """Registers Datasets on ADBC DBAPI connections.""" + + @staticmethod + def matches(con: object) -> TypeGuard[dbapi.Connection]: + # Every driver's ``dbapi.connect`` returns this class or a + # subclass of it (e.g. ``adbc_driver_sqlite.dbapi``'s). + return any( + cls.__module__ == "adbc_driver_manager.dbapi" + and cls.__qualname__ == "Connection" + for cls in type(con).__mro__ + ) + + @staticmethod + def register( + con: dbapi.Connection, + name: str, + ds: xr.Dataset, + *, + chunks: Chunks = None, + table_names: TableNames = None, + mode: IngestMode = "create", + temporary: bool = False, + **kwargs: Any, + ) -> dbapi.Connection: + """Ingest ``ds`` into tables on an ADBC connection. + + Datasets whose variables all share the same dimensions become a + single table named ``name``. Mixed-dimension datasets are split + into one table per dimension group, created in a database schema + named ``name`` so they are queried as ``name.group`` — the same + spelling every other engine uses:: + + xql.register(con, 'era5', ds, table_names={ + ('time', 'latitude', 'longitude'): 'surface', + ('time', 'level', 'latitude', 'longitude'): 'atmosphere', + }) + cur.execute('SELECT ... FROM era5.surface') + + On databases without schemas (SQLite), and for temporary tables, + which most drivers cannot place in a schema, the groups are + created as flat ``name_group`` tables instead. + + Registration runs inside the connection's current transaction: + the tables are visible to this connection immediately, and to + others once you call ``con.commit()`` (unless the connection is + in autocommit mode). + + Args: + mode: What to do when a table already exists, as in ADBC's + ``adbc_ingest``: ``"create"`` (default) raises, + ``"replace"`` drops and recreates it, ``"append"`` and + ``"create_append"`` add rows to it. + temporary: Create temporary tables, which the database drops + when the connection closes. + **kwargs: Forwarded to + [XarrayPushdownDataset][xarray_sql.backends.pyarrow.XarrayPushdownDataset] + (``batch_size``, ``prefetch``, ``prefetch_bytes``, + ``coalesce_rows``) to tune the ingest scan. + """ + groups = group_vars_by_dims(ds) + names = resolve_table_names(ds, table_names, case_insensitive=True) + if len(groups) <= 1: + _ingest( + con, + name, + ds, + chunks, + mode=mode, + temporary=temporary, + **kwargs, + ) + return con + + in_schema = not temporary and _create_schema(con, name) + coord_arrays = shared_coord_arrays(ds) + for dims, var_names in groups.items(): + group = names[dims] + _ingest( + con, + group if in_schema else f"{name}_{group}", + ds[var_names], + chunks, + mode=mode, + temporary=temporary, + db_schema_name=name if in_schema else None, + coord_arrays=coord_arrays, + **kwargs, + ) + return con diff --git a/xarray_sql/backends/base.py b/xarray_sql/backends/base.py index 1446bef8..0fd7ede3 100644 --- a/xarray_sql/backends/base.py +++ b/xarray_sql/backends/base.py @@ -65,7 +65,8 @@ def get_adapter(con: object) -> type[EngineAdapter[Any]]: raise TypeError( f"No xarray-sql engine adapter for connection of type " f"{type(con).__module__}.{type(con).__qualname__}. " - f"Supported: DataFusion SessionContext and DuckDB connections." + f"Supported: DataFusion SessionContext, DuckDB, and ADBC DBAPI " + "connections." ) @@ -108,8 +109,10 @@ def register( Args: con: An engine connection: a ``datafusion.SessionContext`` (or - [xarray_sql.XarrayContext][]) or a - ``duckdb.DuckDBPyConnection``. + [xarray_sql.XarrayContext][]), a + ``duckdb.DuckDBPyConnection``, or an ADBC DBAPI connection + (``adbc_driver_manager.dbapi.Connection``), into which the + Dataset is ingested as a table. name: The table name to register the Dataset under. Datasets whose variables have differing dimensions are split into one table per dimension group, addressed as ``name.group`` on From 793d5d43756df8413e935498d9fd7fd2e79252f5 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 12:28:19 -0700 Subject: [PATCH 02/31] Don't fall back to flat tables through an aborted transaction On PostgreSQL a failed CREATE SCHEMA aborts the transaction (ADBC connections default to autocommit off), so the flat-table fallback's ingests all failed with "connection is in an error state". A savepoint cannot recover it either: the ADBC PostgreSQL driver refuses every statement until rollback. - Use an existing schema as is. PostgreSQL checks the database-level CREATE privilege even for CREATE SCHEMA IF NOT EXISTS, so a role granted only a pre-created schema could not register into it. - After a failed CREATE SCHEMA, fall back to flat names only if the connection still runs statements; otherwise raise a RuntimeError that says to roll back and how to avoid it. - PostgreSQL tests run when XARRAY_SQL_TEST_POSTGRES_URI is set. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- docs/engines.md | 11 +++++-- tests/test_adbc_backend.py | 34 +++++++++++++++++++++ xarray_sql/backends/adbc.py | 60 +++++++++++++++++++++++++++++++++---- 3 files changed, 96 insertions(+), 9 deletions(-) diff --git a/docs/engines.md b/docs/engines.md index d77fc344..94825065 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -246,9 +246,14 @@ Options specific to this adapter: Mixed-dimension Datasets are ingested into a database schema named after the Dataset, so `era5.surface` is the same SQL here as on -DataFusion and DuckDB. On databases without schemas (SQLite), and for -temporary tables, the groups are created as flat `era5_surface` tables -instead (with a warning in the first case). +DataFusion and DuckDB. An existing schema is used as is, so a role +granted only that schema can register into it. On databases without +schemas (SQLite), and for temporary tables, the groups are created as +flat `era5_surface` tables instead (with a warning in the first case). +Where creating the schema fails *and* the failure aborts the +transaction (PostgreSQL without the `CREATE` privilege), registration +raises instead: call `con.rollback()`, then create the schema +beforehand or pass `temporary=True`. The cursor is a one-shot Arrow stream: `xql.to_dataset(cur, ...)` round-trips eagerly, and `chunks=` needs `spill=True`. diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 3ee7497d..04501d83 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -8,6 +8,7 @@ """ import importlib.util +import os import numpy as np import pandas as pd @@ -223,3 +224,36 @@ def test_temporary_mixed_dimensions_use_flat_names(duckdb_con, mixed_ds): count = _query(duckdb_con, "SELECT COUNT(*) FROM era5_surface") assert count.fetchone()[0] == 6 * 3 * 4 + + +@pytest.fixture +def postgres_con(): + uri = os.environ.get("XARRAY_SQL_TEST_POSTGRES_URI") + if not uri: + pytest.skip( + "set XARRAY_SQL_TEST_POSTGRES_URI to run against PostgreSQL" + ) + postgres = pytest.importorskip("adbc_driver_postgresql.dbapi") + connection = postgres.connect(uri) + yield connection + connection.rollback() + connection.close() + + +def test_postgres_schema_failure_explains_the_aborted_transaction( + postgres_con, mixed_ds +): + # PostgreSQL rejects schema names starting with `pg_`, so CREATE SCHEMA + # fails here for any user, and a failed statement aborts the + # transaction: no fallback ingest could run after it. + with pytest.raises(RuntimeError, match="rollback"): + xql.register(postgres_con, "pg_era5", mixed_ds, table_names=NAMES) + + +def test_postgres_uses_an_existing_schema(postgres_con, mixed_ds): + with postgres_con.cursor() as cur: + cur.execute("CREATE SCHEMA IF NOT EXISTS era5") + xql.register(postgres_con, "era5", mixed_ds, table_names=NAMES) + + count = _query(postgres_con, "SELECT COUNT(*) FROM era5.surface") + assert count.fetchone()[0] == 6 * 3 * 4 diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index 4c826c74..559eeb12 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -86,16 +86,60 @@ def _ingest( ) +def _schema_exists(con: dbapi.Connection, name: str) -> bool: + """Whether the database already has a schema named exactly *name*.""" + try: + objects = con.adbc_get_objects( + depth="db_schemas", db_schema_filter=name + ).read_all() + except Exception: # noqa: BLE001 — metadata unsupported; assume not + return False + return any( + schema["db_schema_name"] == name + for catalog in objects.to_pylist() + for schema in catalog["catalog_db_schemas"] or [] + ) + + +def _connection_usable(con: dbapi.Connection) -> bool: + """Whether *con* still runs statements after one of them failed.""" + try: + with con.cursor() as cur: + cur.execute("SELECT 1") + cur.fetchall() + except Exception: # noqa: BLE001 + return False + return True + + def _create_schema(con: dbapi.Connection, name: str) -> bool: - """Create the database schema *name*; whether that succeeded. + """Ensure the database schema *name* exists; whether it does. + + An existing schema is used as is: creating it can need privileges on + the whole database (PostgreSQL checks them even for ``IF NOT + EXISTS``) that a role granted only that schema lacks. Not every ADBC database has schemas (SQLite does not), and creating - one can need privileges the connection lacks. + one can fail for lack of privileges. When the connection survives the + failure, the caller falls back to flat table names. On databases + where a failed statement aborts the transaction (PostgreSQL), nothing + after it could run, so this raises instead of falling back. """ + if _schema_exists(con, name): + return True try: with con.cursor() as cur: cur.execute(f"CREATE SCHEMA IF NOT EXISTS {_quote(name)}") - except Exception as exc: # noqa: BLE001 — fall back to flat names + except Exception as exc: + if not _connection_usable(con): + raise RuntimeError( + f"Could not create the {name!r} schema to hold the dimension " + f"groups of {name!r} ({exc}). The failure aborted the " + f"connection's transaction, so call con.rollback() before " + f"using it again. Create the schema beforehand (or grant " + f"the privilege to), or pass temporary=True to register " + f"flat {name}_ tables instead." + ) from exc warnings.warn( f"Could not create the {name!r} schema to hold the dimension " f"groups of {name!r} ({exc}); registering them as flat " @@ -147,9 +191,13 @@ def register( }) cur.execute('SELECT ... FROM era5.surface') - On databases without schemas (SQLite), and for temporary tables, - which most drivers cannot place in a schema, the groups are - created as flat ``name_group`` tables instead. + An existing schema is used as is. On databases without schemas + (SQLite), and for temporary tables, which most drivers cannot + place in a schema, the groups are created as flat + ``name_group`` tables instead. If creating the schema fails in a + way that aborts the transaction (PostgreSQL without the + ``CREATE`` privilege), this raises ``RuntimeError``; roll back, + then create the schema beforehand or pass ``temporary=True``. Registration runs inside the connection's current transaction: the tables are visible to this connection immediately, and to From 41e94e0d551c218bbe0ffe903d3ffa5d59863e8a Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 12:57:46 -0700 Subject: [PATCH 03/31] Support ClickHouse in the ADBC adapter ClickHouse's ADBC driver only implements append ingest, so the default mode="create" failed with NOT_IMPLEMENTED, and mixed-dimension Datasets could not be registered: ClickHouse has CREATE DATABASE, not CREATE SCHEMA, and the driver implements neither get_info nor get_objects. On ClickHouse the adapter now creates tables itself and appends: - MergeTree tables sorted by their dimensions, so ClickHouse's primary index skips data on dimension predicates; Memory for temporary ones. - create / replace / create_append map onto CREATE, DROP + CREATE, and CREATE ... IF NOT EXISTS. - Timestamps are declared DateTime64(n, 'UTC'), so string literals in queries mean UTC rather than the server's local zone. - Mixed-dimension Datasets go into a ClickHouse database named after the Dataset, queried as name.group like everywhere else. ClickHouse is detected by vendor name, or, for drivers without get_info, by probing system.one. Tests run when XARRAY_SQL_TEST_CLICKHOUSE_URI and XARRAY_SQL_TEST_CLICKHOUSE_DRIVER are set. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 82 +++++++++++++++++++ xarray_sql/backends/adbc.py | 154 ++++++++++++++++++++++++++++++++++-- 2 files changed, 231 insertions(+), 5 deletions(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 04501d83..8ded7a97 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -257,3 +257,85 @@ def test_postgres_uses_an_existing_schema(postgres_con, mixed_ds): count = _query(postgres_con, "SELECT COUNT(*) FROM era5.surface") assert count.fetchone()[0] == 6 * 3 * 4 + + +@pytest.fixture +def clickhouse_con(): + uri = os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_URI") + driver = os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_DRIVER") + if not (uri and driver): + pytest.skip( + "set XARRAY_SQL_TEST_CLICKHOUSE_URI and " + "XARRAY_SQL_TEST_CLICKHOUSE_DRIVER to run against ClickHouse" + ) + connection = dbapi.connect(driver=driver, db_kwargs={"uri": uri}) + for statement in [ + "DROP TABLE IF EXISTS weather", + "DROP DATABASE IF EXISTS era5", + ]: + _query(connection, statement).close() + yield connection + connection.close() + + +def test_clickhouse_round_trips(clickhouse_con, ds): + xql.register(clickhouse_con, "weather", ds) + + cur = _query( + clickhouse_con, + "SELECT time, lat, lon, temperature, precipitation FROM weather " + "ORDER BY time, lat, lon", + ) + xr.testing.assert_allclose(xql.to_dataset(cur, template=ds), ds.compute()) + + +def test_clickhouse_time_literals_mean_utc(clickhouse_con, ds): + xql.register(clickhouse_con, "weather", ds) + + cur = _query( + clickhouse_con, + "SELECT time, lat, lon, temperature FROM weather " + "WHERE time >= '2021-01-01 04:00:00' ORDER BY time, lat, lon", + ) + out = xql.to_dataset(cur, template=ds) + + expected = ds.temperature.isel(time=slice(4, None)) + xr.testing.assert_allclose(out.temperature, expected.compute()) + + +def test_clickhouse_replace_then_append(clickhouse_con, ds): + xql.register(clickhouse_con, "weather", ds) + xql.register( + clickhouse_con, "weather", ds.isel(time=slice(0, 4)), mode="replace" + ) + xql.register( + clickhouse_con, "weather", ds.isel(time=slice(4, 8)), mode="append" + ) + + cur = _query( + clickhouse_con, + "SELECT time, lat, lon, temperature, precipitation FROM weather " + "ORDER BY time, lat, lon", + ) + xr.testing.assert_allclose(xql.to_dataset(cur, template=ds), ds.compute()) + + +def test_clickhouse_temporary_table_is_queryable(clickhouse_con, ds): + xql.register(clickhouse_con, "weather", ds, temporary=True) + + count = _query(clickhouse_con, "SELECT COUNT(*) FROM weather") + assert count.fetchone()[0] == 8 * 5 * 6 + + +def test_clickhouse_mixed_dimensions_register_in_a_database( + clickhouse_con, mixed_ds +): + xql.register(clickhouse_con, "era5", mixed_ds, table_names=NAMES) + + cur = _query( + clickhouse_con, + "SELECT time, level, lat, lon, temperature FROM era5.atmosphere " + "ORDER BY time, level, lat, lon", + ) + out = xql.to_dataset(cur, template=mixed_ds[["temperature"]]) + xr.testing.assert_allclose(out.temperature, mixed_ds.temperature.compute()) diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index 559eeb12..ea0b2c19 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -31,6 +31,7 @@ import warnings from typing import TYPE_CHECKING, Any, Literal, TypeGuard +import pyarrow as pa import xarray as xr from ..df import ( @@ -57,14 +58,127 @@ def _quote(identifier: str) -> str: return f'"{escaped}"' +def _is_clickhouse(con: dbapi.Connection) -> bool: + """Whether *con* is connected to ClickHouse. + + Drivers name their database in ``adbc_get_info``. The ClickHouse + driver does not implement it, so only a driver without it is probed + with a query against ClickHouse's ``system.one`` table. + """ + try: + info = con.adbc_get_info() + except Exception: # noqa: BLE001 — unimplemented; probe instead + pass + else: + return str(info.get("vendor_name", "")).lower() == "clickhouse" + try: + with con.cursor() as cur: + cur.execute("SELECT 1 FROM system.one") + cur.fetchall() + except Exception: # noqa: BLE001 + return False + return True + + +_CLICKHOUSE_TYPES = { + pa.bool_(): "Bool", + pa.int8(): "Int8", + pa.int16(): "Int16", + pa.int32(): "Int32", + pa.int64(): "Int64", + pa.uint8(): "UInt8", + pa.uint16(): "UInt16", + pa.uint32(): "UInt32", + pa.uint64(): "UInt64", + pa.float16(): "Float32", + pa.float32(): "Float32", + pa.float64(): "Float64", + pa.string(): "String", + pa.large_string(): "String", + pa.binary(): "String", + pa.large_binary(): "String", + pa.date32(): "Date32", +} + +_TIMESTAMP_PRECISION = {"s": 0, "ms": 3, "us": 6, "ns": 9} + + +def _clickhouse_type(field: pa.Field, key: bool) -> str: + """The ClickHouse column type for an Arrow field. + + Timestamps without a zone are declared UTC, which is what their + values mean; ClickHouse also parses string literals compared with a + column in that column's zone, so ``time >= '2020-01-01'`` means UTC + rather than the server's local time. Sort-key columns and floats + (which carry NaN) are not ``Nullable``. + """ + arrow_type = field.type + if pa.types.is_timestamp(arrow_type): + precision = _TIMESTAMP_PRECISION[arrow_type.unit] + zone = arrow_type.tz or "UTC" + name = f"DateTime64({precision}, '{zone}')" + elif arrow_type in _CLICKHOUSE_TYPES: + name = _CLICKHOUSE_TYPES[arrow_type] + else: + raise TypeError( + f"no ClickHouse column type for {field.name!r} of Arrow type " + f"{arrow_type}; create the table yourself and register with " + f'mode="append"' + ) + if key or pa.types.is_floating(arrow_type) or not field.nullable: + return name + return f"Nullable({name})" + + +def _clickhouse_ddl( + table: str, + schema: pa.Schema, + dims: tuple[str, ...], + *, + mode: IngestMode, + temporary: bool, + database: str | None, +) -> list[str]: + """Statements that prepare *table* for an append-mode ingest. + + ClickHouse's ADBC driver only appends, so the table is created here + for every other mode. Tables are sorted by their dimensions, so + ClickHouse's primary index skips data on dimension predicates the + way chunk pruning does in the other engines. + """ + if mode == "append": + return [] + target = _quote(table) + if database is not None: + target = f"{_quote(database)}.{target}" + columns = ", ".join( + f"{_quote(field.name)} {_clickhouse_type(field, field.name in dims)}" + for field in schema + ) + kind = "TEMPORARY TABLE" if temporary else "TABLE" + if temporary: + engine = "ENGINE = Memory" + else: + order = ", ".join(_quote(dim) for dim in dims) or "tuple()" + engine = f"ENGINE = MergeTree ORDER BY ({order})" + statements = [] + if mode == "replace": + statements.append(f"DROP {kind} IF EXISTS {target}") + exists = " IF NOT EXISTS" if mode == "create_append" else "" + statements.append(f"CREATE {kind}{exists} {target} ({columns}) {engine}") + return statements + + def _ingest( con: dbapi.Connection, table: str, ds: xr.Dataset, chunks: Chunks, *, + dims: tuple[str, ...], mode: IngestMode, temporary: bool, + clickhouse: bool, db_schema_name: str | None = None, **kwargs: Any, ) -> None: @@ -75,7 +189,20 @@ def _ingest( prefetches chunks on a thread pool while the driver writes earlier batches, so the source read and the database write overlap. """ - reader = XarrayPushdownDataset(ds, chunks, **kwargs).scanner().to_reader() + dataset = XarrayPushdownDataset(ds, chunks, **kwargs) + if clickhouse: + for statement in _clickhouse_ddl( + table, + dataset.schema, + dims, + mode=mode, + temporary=temporary, + database=db_schema_name, + ): + with con.cursor() as cur: + cur.execute(statement) + mode, temporary = "append", False + reader = dataset.scanner().to_reader() with con.cursor() as cur: cur.adbc_ingest( table, @@ -112,7 +239,9 @@ def _connection_usable(con: dbapi.Connection) -> bool: return True -def _create_schema(con: dbapi.Connection, name: str) -> bool: +def _create_schema( + con: dbapi.Connection, name: str, *, clickhouse: bool = False +) -> bool: """Ensure the database schema *name* exists; whether it does. An existing schema is used as is: creating it can need privileges on @@ -123,13 +252,15 @@ def _create_schema(con: dbapi.Connection, name: str) -> bool: one can fail for lack of privileges. When the connection survives the failure, the caller falls back to flat table names. On databases where a failed statement aborts the transaction (PostgreSQL), nothing - after it could run, so this raises instead of falling back. + after it could run, so this raises instead of falling back. In + ClickHouse, a database plays the role of a schema. """ if _schema_exists(con, name): return True + kind = "DATABASE" if clickhouse else "SCHEMA" try: with con.cursor() as cur: - cur.execute(f"CREATE SCHEMA IF NOT EXISTS {_quote(name)}") + cur.execute(f"CREATE {kind} IF NOT EXISTS {_quote(name)}") except Exception as exc: if not _connection_usable(con): raise RuntimeError( @@ -199,6 +330,12 @@ def register( ``CREATE`` privilege), this raises ``RuntimeError``; roll back, then create the schema beforehand or pass ``temporary=True``. + On ClickHouse, whose driver can only append, the adapter creates + the tables itself: ``MergeTree`` tables sorted by their + dimensions (``Memory`` for temporary ones), in a ClickHouse + database named ``name`` for a mixed-dimension Dataset, with + timestamps declared ``DateTime64(9, 'UTC')``. + Registration runs inside the connection's current transaction: the tables are visible to this connection immediately, and to others once you call ``con.commit()`` (unless the connection is @@ -218,19 +355,24 @@ def register( """ groups = group_vars_by_dims(ds) names = resolve_table_names(ds, table_names, case_insensitive=True) + clickhouse = _is_clickhouse(con) if len(groups) <= 1: _ingest( con, name, ds, chunks, + dims=next(iter(groups), ()), mode=mode, temporary=temporary, + clickhouse=clickhouse, **kwargs, ) return con - in_schema = not temporary and _create_schema(con, name) + in_schema = not temporary and _create_schema( + con, name, clickhouse=clickhouse + ) coord_arrays = shared_coord_arrays(ds) for dims, var_names in groups.items(): group = names[dims] @@ -239,8 +381,10 @@ def register( group if in_schema else f"{name}_{group}", ds[var_names], chunks, + dims=dims, mode=mode, temporary=temporary, + clickhouse=clickhouse, db_schema_name=name if in_schema else None, coord_arrays=coord_arrays, **kwargs, From 8d57591df6af589992215a5c1ef0681ac93d2f72 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 13:17:23 -0700 Subject: [PATCH 04/31] Document ClickHouse in the ADBC section Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- docs/engines.md | 24 +++++++++++++++++++++++- 1 file changed, 23 insertions(+), 1 deletion(-) diff --git a/docs/engines.md b/docs/engines.md index 94825065..3cba2475 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -255,6 +255,28 @@ transaction (PostgreSQL without the `CREATE` privilege), registration raises instead: call `con.rollback()`, then create the schema beforehand or pass `temporary=True`. +**ClickHouse.** ClickHouse's +[ADBC driver](https://adbc-drivers.org/drivers/clickhouse/) (a preview +at the time of writing) can only append, so on ClickHouse the adapter +creates each table itself and then appends to it: + +```python +from adbc_driver_manager import dbapi + +con = dbapi.connect(driver="clickhouse", db_kwargs={"uri": "http://localhost:8123/"}) +xql.register(con, "era5", ds) +``` + +Tables are `MergeTree` sorted by their dimensions +(`ORDER BY (time, latitude, longitude)`), so ClickHouse's primary index +skips data on dimension filters much as chunk pruning does elsewhere. +Timestamps are declared `DateTime64(9, 'UTC')`, so a literal like +`time >= '2020-01-01'` means UTC rather than the server's local zone. +Mixed-dimension Datasets go into a ClickHouse *database* named after +the Dataset (`era5.surface`), and `temporary=True` creates `Memory` +tables. To choose the engine or sort key yourself, create the table +first and register with `mode="append"`. + The cursor is a one-shot Arrow stream: `xql.to_dataset(cur, ...)` round-trips eagerly, and `chunks=` needs `spill=True`. @@ -273,7 +295,7 @@ What each integration provides. Known issues and constraints live on | `geometry` column ([geospatial](geospatial.md#geoarrow-point-geometry-columns)) | annotated WKB passes through | native `GEOMETRY` (`"wkb"` encoding) | plain binary/struct | driver-dependent | | Mixed-dimension datasets | one schema, `name.group` tables | `name.group` views over `name_group` tables | `xql.arrow_datasets(ds, name)`, one per group | `name.group` tables in a schema; `name_group` without schemas | | Naming those tables (`table_names=`) | yes | yes | yes | yes | -| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.0` (tested on 1.12 with SQLite and DuckDB drivers) | +| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.0` (tested on 1.12 with SQLite, DuckDB, PostgreSQL 18, and ClickHouse 26.8 drivers) | [^spill-only]: Why DuckDB relations do not re-execute — and two other engine-specific issues worth knowing — is explained on From cac9fc67a68b60527670dd3500dd7e11aeba7331 Mon Sep 17 00:00:00 2001 From: Stephen Kent <43362477+kentstephen@users.noreply.github.com> Date: Sat, 26 Sep 2026 20:01:55 -0400 Subject: [PATCH 05/31] Keep missing values in ClickHouse float columns (#257) --- tests/test_adbc_backend.py | 43 +++++++++++++++++++++++++++++++++++++ xarray_sql/backends/adbc.py | 8 ++++--- 2 files changed, 48 insertions(+), 3 deletions(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 8ded7a97..8671b015 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -12,10 +12,12 @@ import numpy as np import pandas as pd +import pyarrow as pa import pytest import xarray as xr import xarray_sql as xql +from xarray_sql.backends.adbc import _clickhouse_ddl dbapi = pytest.importorskip("adbc_driver_manager.dbapi") sqlite_dbapi = pytest.importorskip("adbc_driver_sqlite.dbapi") @@ -226,6 +228,30 @@ def test_temporary_mixed_dimensions_use_flat_names(duckdb_con, mixed_ds): assert count.fetchone()[0] == 6 * 3 * 4 +def test_clickhouse_float_columns_are_nullable(): + # The scan writes NaN as an Arrow null, so aggregates skip it; a plain + # Float64 column would store that null as 0. + schema = pa.schema( + [ + pa.field("time", pa.timestamp("ns")), + pa.field("lat", pa.float64()), + pa.field("t2m", pa.float64()), + pa.field("sst", pa.float32()), + ] + ) + [ddl] = _clickhouse_ddl( + "weather", + schema, + ("time", "lat"), + mode="create", + temporary=False, + database=None, + ) + assert '"t2m" Nullable(Float64)' in ddl + assert '"sst" Nullable(Float32)' in ddl + assert '"lat" Float64,' in ddl # sort keys stay non-Nullable + + @pytest.fixture def postgres_con(): uri = os.environ.get("XARRAY_SQL_TEST_POSTGRES_URI") @@ -339,3 +365,20 @@ def test_clickhouse_mixed_dimensions_register_in_a_database( ) out = xql.to_dataset(cur, template=mixed_ds[["temperature"]]) xr.testing.assert_allclose(out.temperature, mixed_ds.temperature.compute()) + + +def test_clickhouse_keeps_missing_values(clickhouse_con, ds): + holed = ds.copy(deep=True) + holed["temperature"][0, 0, 0] = np.nan + xql.register(clickhouse_con, "weather", holed) + + with _query(clickhouse_con, "SELECT AVG(temperature) FROM weather") as avg: + mean = avg.fetchone()[0] + assert mean == pytest.approx(float(holed.temperature.mean())) + cur = _query( + clickhouse_con, + "SELECT time, lat, lon, temperature FROM weather " + "ORDER BY time, lat, lon", + ) + out = xql.to_dataset(cur, template=holed) + xr.testing.assert_allclose(out.temperature, holed.temperature.compute()) diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index ea0b2c19..d7b665d8 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -109,8 +109,10 @@ def _clickhouse_type(field: pa.Field, key: bool) -> str: Timestamps without a zone are declared UTC, which is what their values mean; ClickHouse also parses string literals compared with a column in that column's zone, so ``time >= '2020-01-01'`` means UTC - rather than the server's local time. Sort-key columns and floats - (which carry NaN) are not ``Nullable``. + rather than the server's local time. Sort-key columns are not + ``Nullable``. Floats are: the scan writes NaN as null so aggregates + skip missing values, and a plain ``Float64`` column would store that + null as 0. """ arrow_type = field.type if pa.types.is_timestamp(arrow_type): @@ -125,7 +127,7 @@ def _clickhouse_type(field: pa.Field, key: bool) -> str: f"{arrow_type}; create the table yourself and register with " f'mode="append"' ) - if key or pa.types.is_floating(arrow_type) or not field.nullable: + if key or not field.nullable: return name return f"Nullable({name})" From 82a2c6d945ac5d5946f2f50212c054d733897996 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 17:09:44 -0700 Subject: [PATCH 06/31] Compare zone-aware result times with zone-labeled window literals The chunked round-trip builds window predicates from numpy times, which carry no zone. Polars refuses to compare those with a zone-aware column, so a result whose times are labeled (ClickHouse's DateTime64(n, 'UTC'), PostgreSQL's timestamptz, a pyarrow table with tz=...) failed with "could not evaluate '<' comparison" on every spilled or Polars chunked reconstruction; the eager path was unaffected. Every spill reads through PolarsHandle, so the fix is there: a naive window value is a UTC instant, so it is labeled UTC and expressed in the column's zone. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_lazy_roundtrip.py | 26 ++++++++++++++++++++++++++ xarray_sql/lazyscan.py | 27 ++++++++++++++++++++++++--- 2 files changed, 50 insertions(+), 3 deletions(-) diff --git a/tests/test_lazy_roundtrip.py b/tests/test_lazy_roundtrip.py index 185191ef..f220be81 100644 --- a/tests/test_lazy_roundtrip.py +++ b/tests/test_lazy_roundtrip.py @@ -375,3 +375,29 @@ def test_polars_spill_uses_streaming_sink(source, tmp_path): lf, template=source, chunks={"time": 20}, spill=tmp_path ) xr.testing.assert_allclose(out.compute(), source) + + +@pytest.mark.parametrize("zone", ["UTC", "Asia/Tokyo"]) +def test_chunked_round_trip_of_zone_aware_times(zone): + # Databases such as ClickHouse and PostgreSQL return times labeled + # with a zone; the template's numpy times are plain UTC instants. + ds = xr.Dataset( + {"temperature": (["time", "lat"], np.random.rand(8, 3))}, + coords={ + "time": pd.date_range("2021-01-01", periods=8, freq="h"), + "lat": [1.0, 2.0, 3.0], + }, + ) + result = pa.table( + { + "time": pa.array(np.repeat(ds.time.values, 3)).cast( + pa.timestamp("us", tz=zone) + ), + "lat": np.tile(ds.lat.values, 8), + "temperature": ds.temperature.values.ravel(), + } + ) + + out = xql.to_dataset(result, template=ds, chunks={"time": 2}, spill=True) + + xr.testing.assert_allclose(out.compute(), ds) diff --git a/xarray_sql/lazyscan.py b/xarray_sql/lazyscan.py index 1f4b71c8..91de0a2e 100644 --- a/xarray_sql/lazyscan.py +++ b/xarray_sql/lazyscan.py @@ -75,6 +75,20 @@ def _plain(value: Any) -> Any: return value +def _zoned(value: Any, zone: str | None) -> Any: + """A plain literal comparable with a column in time zone *zone*. + + Window values come from numpy, which has no time zones: a + ``datetime64`` is a UTC instant. Engines refuse to compare a naive + literal with a zone-aware column (or compare it as local time), so + the literal is labeled UTC and expressed in the column's zone. + """ + plain = _plain(value) + if zone and isinstance(plain, pd.Timestamp) and plain.tzinfo is None: + return plain.tz_localize("UTC").tz_convert(zone) + return plain + + class LazyResultHandle(Protocol): """A re-executable query result (see module docstring).""" @@ -300,10 +314,14 @@ def fetch( ) -> list[pa.RecordBatch]: import polars as pl + schema = self._lf.collect_schema() exprs = [] for dim, (kind, a, b) in specs.items(): + zone = getattr(schema.get(dim), "time_zone", None) if kind == "range": - exprs.append(pl.col(dim).is_between(_plain(a), _plain(b))) + exprs.append( + pl.col(dim).is_between(_zoned(a, zone), _zoned(b, zone)) + ) elif getattr(a, "dtype", None) is not None and a.dtype.kind == "f": # Upstream Polars translates float ``is_in`` literals # imprecisely (silently matching nothing); degenerate @@ -313,11 +331,14 @@ def fetch( # of values. exprs.append( pl.any_horizontal( - [pl.col(dim).is_between(*(_plain(v),) * 2) for v in a] + [ + pl.col(dim).is_between(*(_zoned(v, zone),) * 2) + for v in a + ] ) ) else: - exprs.append(pl.col(dim).is_in([_plain(v) for v in a])) + exprs.append(pl.col(dim).is_in([_zoned(v, zone) for v in a])) lf = self._lf.filter(*exprs) if exprs else self._lf out = _collect_streaming(lf.select([pl.col(n) for n in columns])) return cast(list[pa.RecordBatch], out.to_arrow().to_batches()) From af18619dbf903a2bb51e28d9b7f2800bdd59e64d Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 17:13:25 -0700 Subject: [PATCH 07/31] Quote identifiers in each database's own SQL; register on MySQL CREATE SCHEMA "era5" is a syntax error on MySQL (double quotes delimit strings there unless ANSI_QUOTES is set, and always on BigQuery), so mixed-dimension Datasets fell back to flat era5_ tables and era5.group was not portable. The adapter now looks up the vendor once per register and quotes with backticks for MySQL, MariaDB, and BigQuery. With the schema created, the MySQL driver (0.6.1) then failed the ingest: given a target schema, it creates the table in the connection's default database but inserts into the target. For MySQL the adapter makes the target the default database for the ingest instead, and restores the previous default afterwards. MySQL tests run when XARRAY_SQL_TEST_MYSQL_URI is set. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 40 +++++++++++++++++ xarray_sql/backends/adbc.py | 85 +++++++++++++++++++++++++++---------- 2 files changed, 103 insertions(+), 22 deletions(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 8671b015..26befe82 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -382,3 +382,43 @@ def test_clickhouse_keeps_missing_values(clickhouse_con, ds): ) out = xql.to_dataset(cur, template=holed) xr.testing.assert_allclose(out.temperature, holed.temperature.compute()) + + +@pytest.fixture +def mysql_con(): + uri = os.environ.get("XARRAY_SQL_TEST_MYSQL_URI") + if not uri: + pytest.skip("set XARRAY_SQL_TEST_MYSQL_URI to run against MySQL") + connection = dbapi.connect(driver="mysql", db_kwargs={"uri": uri}) + for statement in [ + "DROP TABLE IF EXISTS weather", + "DROP DATABASE IF EXISTS era5", + ]: + _query(connection, statement).close() + yield connection + connection.close() + + +def test_mysql_round_trips(mysql_con, ds): + xql.register(mysql_con, "weather", ds) + + cur = _query( + mysql_con, + "SELECT time, lat, lon, temperature, precipitation FROM weather " + "ORDER BY time, lat, lon", + ) + xr.testing.assert_allclose(xql.to_dataset(cur, template=ds), ds.compute()) + + +def test_mysql_mixed_dimensions_register_in_a_database(mysql_con, mixed_ds): + xql.register(mysql_con, "era5", mixed_ds, table_names=NAMES) + + cur = _query( + mysql_con, + "SELECT time, level, lat, lon, temperature FROM era5.atmosphere " + "ORDER BY time, level, lat, lon", + ) + out = xql.to_dataset(cur, template=mixed_ds[["temperature"]]) + xr.testing.assert_allclose(out.temperature, mixed_ds.temperature.compute()) + default = _query(mysql_con, "SELECT DATABASE()").fetchone()[0] + assert default != "era5" diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index d7b665d8..63655a18 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -28,7 +28,9 @@ from __future__ import annotations +import contextlib import warnings +from collections.abc import Iterator from typing import TYPE_CHECKING, Any, Literal, TypeGuard import pyarrow as pa @@ -52,14 +54,25 @@ IngestMode = Literal["create", "append", "replace", "create_append"] -def _quote(identifier: str) -> str: - """Render *identifier* as a quoted SQL identifier.""" +_BACKTICK_VENDORS = ("mysql", "mariadb", "bigquery") +"""Databases whose SQL quotes identifiers with backticks, not ``"``. + +MySQL reads ``"era5"`` as a string unless ``ANSI_QUOTES`` is set, and +BigQuery always does. +""" + + +def _quote(identifier: str, vendor: str = "") -> str: + """Render *identifier* as a quoted SQL identifier in *vendor*'s SQL.""" + if any(name in vendor for name in _BACKTICK_VENDORS): + escaped = identifier.replace("`", "``") + return f"`{escaped}`" escaped = identifier.replace('"', '""') return f'"{escaped}"' -def _is_clickhouse(con: dbapi.Connection) -> bool: - """Whether *con* is connected to ClickHouse. +def _vendor(con: dbapi.Connection) -> str: + """The lowercase name of the database behind *con*; ``""`` if unknown. Drivers name their database in ``adbc_get_info``. The ClickHouse driver does not implement it, so only a driver without it is probed @@ -70,14 +83,14 @@ def _is_clickhouse(con: dbapi.Connection) -> bool: except Exception: # noqa: BLE001 — unimplemented; probe instead pass else: - return str(info.get("vendor_name", "")).lower() == "clickhouse" + return str(info.get("vendor_name") or "").lower() try: with con.cursor() as cur: cur.execute("SELECT 1 FROM system.one") cur.fetchall() except Exception: # noqa: BLE001 - return False - return True + return "" + return "clickhouse" _CLICKHOUSE_TYPES = { @@ -180,7 +193,7 @@ def _ingest( dims: tuple[str, ...], mode: IngestMode, temporary: bool, - clickhouse: bool, + vendor: str, db_schema_name: str | None = None, **kwargs: Any, ) -> None: @@ -192,7 +205,7 @@ def _ingest( batches, so the source read and the database write overlap. """ dataset = XarrayPushdownDataset(ds, chunks, **kwargs) - if clickhouse: + if vendor == "clickhouse": for statement in _clickhouse_ddl( table, dataset.schema, @@ -205,16 +218,47 @@ def _ingest( cur.execute(statement) mode, temporary = "append", False reader = dataset.scanner().to_reader() - with con.cursor() as cur: + target_schema = db_schema_name + in_database: contextlib.AbstractContextManager[None] = ( + contextlib.nullcontext() + ) + if db_schema_name is not None and vendor.startswith(("mysql", "mariadb")): + # The MySQL driver creates the table in the connection's default + # database but inserts into the one it was given, so name the + # target by making it the default instead. + in_database = _default_database(con, db_schema_name, vendor) + target_schema = None + with in_database, con.cursor() as cur: cur.adbc_ingest( table, reader, mode=mode, - db_schema_name=db_schema_name, + db_schema_name=target_schema, temporary=temporary, ) +@contextlib.contextmanager +def _default_database( + con: dbapi.Connection, database: str, vendor: str +) -> Iterator[None]: + """Make *database* the MySQL connection's default while in the block. + + The previous default is restored afterwards. MySQL cannot unset a + default database, so a connection that had none keeps *database*. + """ + with con.cursor() as cur: + cur.execute("SELECT DATABASE()") + (previous,) = cur.fetchone() + cur.execute(f"USE {_quote(database, vendor)}") + try: + yield + finally: + if previous is not None: + with con.cursor() as cur: + cur.execute(f"USE {_quote(previous, vendor)}") + + def _schema_exists(con: dbapi.Connection, name: str) -> bool: """Whether the database already has a schema named exactly *name*.""" try: @@ -241,9 +285,7 @@ def _connection_usable(con: dbapi.Connection) -> bool: return True -def _create_schema( - con: dbapi.Connection, name: str, *, clickhouse: bool = False -) -> bool: +def _create_schema(con: dbapi.Connection, name: str, *, vendor: str) -> bool: """Ensure the database schema *name* exists; whether it does. An existing schema is used as is: creating it can need privileges on @@ -259,10 +301,11 @@ def _create_schema( """ if _schema_exists(con, name): return True - kind = "DATABASE" if clickhouse else "SCHEMA" + kind = "DATABASE" if vendor == "clickhouse" else "SCHEMA" + target = _quote(name, vendor) try: with con.cursor() as cur: - cur.execute(f"CREATE {kind} IF NOT EXISTS {_quote(name)}") + cur.execute(f"CREATE {kind} IF NOT EXISTS {target}") except Exception as exc: if not _connection_usable(con): raise RuntimeError( @@ -357,7 +400,7 @@ def register( """ groups = group_vars_by_dims(ds) names = resolve_table_names(ds, table_names, case_insensitive=True) - clickhouse = _is_clickhouse(con) + vendor = _vendor(con) if len(groups) <= 1: _ingest( con, @@ -367,14 +410,12 @@ def register( dims=next(iter(groups), ()), mode=mode, temporary=temporary, - clickhouse=clickhouse, + vendor=vendor, **kwargs, ) return con - in_schema = not temporary and _create_schema( - con, name, clickhouse=clickhouse - ) + in_schema = not temporary and _create_schema(con, name, vendor=vendor) coord_arrays = shared_coord_arrays(ds) for dims, var_names in groups.items(): group = names[dims] @@ -386,7 +427,7 @@ def register( dims=dims, mode=mode, temporary=temporary, - clickhouse=clickhouse, + vendor=vendor, db_schema_name=name if in_schema else None, coord_arrays=coord_arrays, **kwargs, From 9ab110a5d21e6a46eba6a87d1a2c5efda1dbd071 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 17:16:30 -0700 Subject: [PATCH 08/31] Round-trip timedelta coordinates through every ADBC database A timedelta64 coordinate (a forecast step) never came back: - SQLite and ClickHouse have no duration type, so ingest failed. Durations are now stored there as integer counts of their unit, which to_dataset reads back as the template's timedelta64 (the Arrow unit is derived from it, so the counts match). - DuckDB and PostgreSQL return month-day-nano intervals and MySQL returns text, which the template dtype cast skipped, leaving an object coordinate. These are now parsed into timedeltas (intervals spanning calendar months, which have no fixed length, are rejected). - The chunked round-trip's spill could not write intervals to Parquet and built windows against text; interval and text dimension columns whose template coordinate is a timedelta are written as durations. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 22 +++++++++++++ tests/test_lazy_roundtrip.py | 18 +++++++++++ xarray_sql/backends/adbc.py | 34 +++++++++++++++++-- xarray_sql/ds.py | 49 +++++++++++++++++++++++++++- xarray_sql/roundtrip.py | 63 ++++++++++++++++++++++++++++++++++-- 5 files changed, 179 insertions(+), 7 deletions(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 26befe82..ba9bbe42 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -161,6 +161,28 @@ def test_chunked_round_trip_spills_the_cursor(duckdb_con, ds): ) +@pytest.mark.parametrize("chunks", [None, {"step": 2}]) +def test_timedelta_coordinates_round_trip(con, chunks): + # Forecast `step`: SQLite stores it as an integer count, DuckDB + # returns it as an interval. + forecast = xr.Dataset( + {"t2m": (["step", "lat"], np.random.rand(4, 2))}, + coords={ + "step": pd.to_timedelta([0, 6, 12, 18], unit="h"), + "lat": [1.0, 2.0], + }, + ).chunk({"step": 2}) + xql.register(con, "forecast", forecast) + + cur = _query(con, "SELECT step, lat, t2m FROM forecast ORDER BY step, lat") + out = xql.to_dataset( + cur, template=forecast, chunks=chunks, spill=chunks is not None + ) + + xr.testing.assert_identical(out.compute().step, forecast.step) + xr.testing.assert_allclose(out.compute(), forecast.compute()) + + def test_existing_table_is_not_overwritten_by_default(con, ds): xql.register(con, "weather", ds) diff --git a/tests/test_lazy_roundtrip.py b/tests/test_lazy_roundtrip.py index f220be81..d8148a58 100644 --- a/tests/test_lazy_roundtrip.py +++ b/tests/test_lazy_roundtrip.py @@ -401,3 +401,21 @@ def test_chunked_round_trip_of_zone_aware_times(zone): out = xql.to_dataset(result, template=ds, chunks={"time": 2}, spill=True) xr.testing.assert_allclose(out.compute(), ds) + + +@pytest.mark.parametrize("chunks", [None, {"step": 2}]) +def test_timedelta_coordinates_returned_as_text(chunks): + # MySQL and Trino have no duration type and return text. + ds = xr.Dataset( + {"t2m": (["step"], np.arange(4.0))}, + coords={"step": pd.to_timedelta([0, 6, 12, 18], unit="h")}, + ) + result = pa.table( + {"step": ["0s", "21600s", "43200s", "64800s"], "t2m": np.arange(4.0)} + ) + + out = xql.to_dataset( + result, template=ds, chunks=chunks, spill=chunks is not None + ) + + xr.testing.assert_identical(out.compute(), ds) diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index 63655a18..0f316c02 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -184,6 +184,33 @@ def _clickhouse_ddl( return statements +_NO_DURATION_VENDORS = ("sqlite", "clickhouse") +"""Databases with no duration type; timedeltas are stored as integers.""" + + +def _durations_as_integers( + reader: pa.RecordBatchReader, +) -> pa.RecordBatchReader: + """*reader* with duration columns as integer counts of their unit. + + The template's ``timedelta64`` unit matches the Arrow unit the scan + derived from it, so [xarray_sql.to_dataset][] reads the counts back + as the original durations. + """ + fields = [ + pa.field(f.name, pa.int64(), f.nullable, f.metadata) + if pa.types.is_duration(f.type) + else f + for f in reader.schema + ] + if all(f.type == g.type for f, g in zip(fields, reader.schema)): + return reader + schema = pa.schema(fields, metadata=reader.schema.metadata) + return pa.RecordBatchReader.from_batches( + schema, (batch.cast(schema) for batch in reader) + ) + + def _ingest( con: dbapi.Connection, table: str, @@ -204,11 +231,13 @@ def _ingest( prefetches chunks on a thread pool while the driver writes earlier batches, so the source read and the database write overlap. """ - dataset = XarrayPushdownDataset(ds, chunks, **kwargs) + reader = XarrayPushdownDataset(ds, chunks, **kwargs).scanner().to_reader() + if vendor in _NO_DURATION_VENDORS: + reader = _durations_as_integers(reader) if vendor == "clickhouse": for statement in _clickhouse_ddl( table, - dataset.schema, + reader.schema, dims, mode=mode, temporary=temporary, @@ -217,7 +246,6 @@ def _ingest( with con.cursor() as cur: cur.execute(statement) mode, temporary = "append", False - reader = dataset.scanner().to_reader() target_schema = db_schema_name in_database: contextlib.AbstractContextManager[None] = ( contextlib.nullcontext() diff --git a/xarray_sql/ds.py b/xarray_sql/ds.py index 94339e0d..0ad1d958 100644 --- a/xarray_sql/ds.py +++ b/xarray_sql/ds.py @@ -70,6 +70,53 @@ def _ds_var_dims(ds: xr.Dataset) -> list[str]: return list(ds.dims) +_TIMEDELTA_PARTS = ( + "weeks", + "days", + "hours", + "minutes", + "seconds", + "milliseconds", + "microseconds", + "nanoseconds", +) + + +def _timedelta(value: Any) -> pd.Timedelta: + """One duration a database returned as an interval or as text. + + Arrow's month-day-nano intervals (DuckDB, PostgreSQL) arrive as + pandas ``DateOffset`` objects; MySQL and Trino return text such as + ``'21600s'``. + """ + if value is None: + return pd.NaT + if isinstance(value, pd.DateOffset): + parts = value.kwds + if parts.get("years") or parts.get("months"): + raise ValueError( + f"{value!r} spans calendar months, which have no fixed duration" + ) + return pd.Timedelta(**{k: parts.get(k, 0) for k in _TIMEDELTA_PARTS}) + return pd.Timedelta(value) + + +def _as_dtype(coord: xr.DataArray, dtype: np.dtype) -> xr.DataArray: + """*coord* cast to the template's *dtype*. + + Databases without a duration type hand timedeltas back as intervals + or text, which numpy cannot cast; those are parsed first. (Integers + need nothing: numpy reads them as counts of the template's unit, + which is how durations are stored where no duration type exists.) + """ + if dtype.kind == "m" and coord.dtype.kind in "OUS": + values = np.array( + [_timedelta(v) for v in coord.values], dtype="timedelta64[ns]" + ) + return coord.copy(data=values.astype(dtype)) + return coord.astype(dtype) + + def _apply_template(ds: xr.Dataset, template: xr.Dataset) -> xr.Dataset: """Recover metadata that the forward SQL pivot strips. @@ -114,7 +161,7 @@ def _apply_template(ds: xr.Dataset, template: xr.Dataset) -> xr.Dataset: tdt = template.coords[d].dtype if out.coords[d].dtype != tdt: try: - out = out.assign_coords({d: out.coords[d].astype(tdt)}) + out = out.assign_coords({d: _as_dtype(out.coords[d], tdt)}) except (ValueError, TypeError): pass # incompatible cast; leave as-is out[d].attrs = dict(template.coords[d].attrs) diff --git a/xarray_sql/roundtrip.py b/xarray_sql/roundtrip.py index e342828f..9e9d93ef 100644 --- a/xarray_sql/roundtrip.py +++ b/xarray_sql/roundtrip.py @@ -29,6 +29,7 @@ from typing import Any, Literal import numpy as np +import pandas as pd import pyarrow as pa import pyarrow.compute as pc import pyarrow.parquet as pq @@ -456,7 +457,7 @@ def _to_dataset_spilled( if handle is not None: handle.spill_parquet(path) else: - _stream_to_parquet(result, path) + _stream_to_parquet(result, path, template) except BaseException: os.unlink(path) raise @@ -481,8 +482,45 @@ def _unlink_quietly(path: str) -> None: pass -def _stream_to_parquet(result: Any, path: str) -> None: - """Write a one-shot Arrow result to Parquet, batch by batch.""" +_MONTH_DAY_NANO = np.dtype( + [("months", " pa.Array: + """Durations a database returned as intervals or text, as ``duration``. + + Databases without a duration type return month-day-nano intervals + (DuckDB, PostgreSQL) or text such as ``'21600s'`` (MySQL, Trino). + """ + if pa.types.is_interval(array.type): + parts = np.frombuffer( + array.buffers()[1], + dtype=_MONTH_DAY_NANO, + count=array.offset + len(array), + )[array.offset :] + valid = ~np.asarray(array.is_null()) + if parts["months"][valid].any(): + raise ValueError( + "an interval spans calendar months, which have no fixed " + "duration" + ) + nanos = parts["days"].astype(np.int64) * _NANOS_PER_DAY + parts["nanos"] + return pa.array(nanos, pa.duration("ns"), mask=~valid) + return pa.array(pd.to_timedelta(array.to_pandas()), pa.duration("ns")) + + +def _stream_to_parquet( + result: Any, path: str, template: xr.Dataset | None = None +) -> None: + """Write a one-shot Arrow result to Parquet, batch by batch. + + Columns the template holds as ``timedelta64`` coordinates but the + database returned as intervals or text are written as durations: + Parquet cannot store month-day-nano intervals, and the windows the + chunked reconstruction builds compare durations, not text. + """ opened = _open_stream(result) if opened is None: raise TypeError( @@ -490,6 +528,25 @@ def _stream_to_parquet(result: Any, path: str) -> None: "Arrow stream." ) schema, batches = opened + durations = [ + i + for i, field in enumerate(schema) + if template is not None + and field.name in template.coords + and template.coords[field.name].dtype.kind == "m" + and ( + pa.types.is_interval(field.type) + or pa.types.is_string(field.type) + or pa.types.is_large_string(field.type) + ) + ] + for i in durations: + schema = schema.set(i, schema.field(i).with_type(pa.duration("ns"))) with pq.ParquetWriter(path, schema) as writer: for batch in batches: + if durations: + columns = list(batch.columns) + for i in durations: + columns[i] = _as_durations(columns[i]) + batch = pa.RecordBatch.from_arrays(columns, schema=schema) writer.write_batch(batch) From 2a68d571cbf513a0c3797484383ff042ae6801c8 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 17:17:03 -0700 Subject: [PATCH 09/31] Document ClickHouse details from review; run ClickHouse tests on chDB in CI - Timestamp precision follows the coordinate's resolution (DateTime64(6, 'UTC') for datetime64[us]), not always 9. - ClickHouse driver 0.1.1 fails against server 26.9; credentials go in URI query parameters, not user-info. - chDB (embedded ClickHouse) takes the same adapter path with no server; CI installs its ADBC driver with dbc and runs the ClickHouse tests against it. - MySQL quoting and duration handling. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- .github/workflows/ci.yml | 12 ++++++++++++ docs/engines.md | 31 +++++++++++++++++++++++++++---- xarray_sql/backends/adbc.py | 3 ++- 3 files changed, 41 insertions(+), 5 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c1d892d9..aafb25d8 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -68,5 +68,17 @@ jobs: run: uv sync --dev --no-install-package xarray-sql - name: build rust run: uv run --no-project maturin develop --uv + - name: Install the chDB ADBC driver + # Embedded ClickHouse: runs the ClickHouse adapter tests without a + # server. dbc installs into the active virtualenv. + env: + VIRTUAL_ENV: ${{ github.workspace }}/.venv + run: | + uv pip install dbc + uv run --no-project dbc install chdb - name: Run unit tests + env: + VIRTUAL_ENV: ${{ github.workspace }}/.venv + XARRAY_SQL_TEST_CLICKHOUSE_DRIVER: chdb + XARRAY_SQL_TEST_CLICKHOUSE_URI: "chdb://" run: uv run --no-project pytest -v . -m "not integration" diff --git a/docs/engines.md b/docs/engines.md index 3cba2475..f359b9b9 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -263,20 +263,43 @@ creates each table itself and then appends to it: ```python from adbc_driver_manager import dbapi -con = dbapi.connect(driver="clickhouse", db_kwargs={"uri": "http://localhost:8123/"}) +# dbc install clickhouse +con = dbapi.connect( + driver="clickhouse", + db_kwargs={"uri": "http://localhost:8123/?user=default&password=..."}, +) xql.register(con, "era5", ds) ``` +Pass credentials as URI query parameters, as above; credentials in the +URI's user-info part are not used. Driver 0.1.1 works with ClickHouse +26.8 but fails every query against 26.9 (`decompression error: incorrect +magic number`). [chDB](https://clickhouse.com/docs/chdb), ClickHouse +embedded in-process, takes the same path with no server: +`dbc install chdb`, then `driver="chdb"` and `uri="chdb://"`. + Tables are `MergeTree` sorted by their dimensions (`ORDER BY (time, latitude, longitude)`), so ClickHouse's primary index skips data on dimension filters much as chunk pruning does elsewhere. -Timestamps are declared `DateTime64(9, 'UTC')`, so a literal like -`time >= '2020-01-01'` means UTC rather than the server's local zone. +Timestamps are declared `DateTime64(p, 'UTC')`, with the precision `p` +following the coordinate's resolution (9 for `datetime64[ns]`, 6 for +`datetime64[us]`), so a literal like `time >= '2020-01-01'` means UTC +rather than the server's local zone. Mixed-dimension Datasets go into a ClickHouse *database* named after the Dataset (`era5.surface`), and `temporary=True` creates `Memory` tables. To choose the engine or sort key yourself, create the table first and register with `mode="append"`. +**MySQL.** Mixed-dimension Datasets go into a MySQL database named +after the Dataset, queried as `era5.surface`; identifiers are quoted +with backticks there (and on BigQuery), as those dialects require. + +**Durations.** A `timedelta64` coordinate (a forecast step, say) +round-trips everywhere, however the database stores it: as a duration +where one exists, as an integer count of its unit on SQLite and +ClickHouse, and as intervals (DuckDB, PostgreSQL) or text (MySQL) that +`to_dataset` converts back using the template. + The cursor is a one-shot Arrow stream: `xql.to_dataset(cur, ...)` round-trips eagerly, and `chunks=` needs `spill=True`. @@ -295,7 +318,7 @@ What each integration provides. Known issues and constraints live on | `geometry` column ([geospatial](geospatial.md#geoarrow-point-geometry-columns)) | annotated WKB passes through | native `GEOMETRY` (`"wkb"` encoding) | plain binary/struct | driver-dependent | | Mixed-dimension datasets | one schema, `name.group` tables | `name.group` views over `name_group` tables | `xql.arrow_datasets(ds, name)`, one per group | `name.group` tables in a schema; `name_group` without schemas | | Naming those tables (`table_names=`) | yes | yes | yes | yes | -| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.0` (tested on 1.12 with SQLite, DuckDB, PostgreSQL 18, and ClickHouse 26.8 drivers) | +| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.0` (tested on 1.12 with SQLite, DuckDB, PostgreSQL 18, MySQL 8.4, ClickHouse 26.8, and chDB 26.7 drivers) | [^spill-only]: Why DuckDB relations do not re-execute — and two other engine-specific issues worth knowing — is explained on diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index 0f316c02..df26e5db 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -407,7 +407,8 @@ def register( the tables itself: ``MergeTree`` tables sorted by their dimensions (``Memory`` for temporary ones), in a ClickHouse database named ``name`` for a mixed-dimension Dataset, with - timestamps declared ``DateTime64(9, 'UTC')``. + timestamps declared ``DateTime64(p, 'UTC')`` at the coordinate's + precision (``p`` is 9 for ``datetime64[ns]``). Registration runs inside the connection's current transaction: the tables are visible to this connection immediately, and to From 3d321440524f82c0db91f97d189955674c4eef39 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 18:18:39 -0700 Subject: [PATCH 10/31] Narrow widened types back to the template; parse text times on spill Databases without a type widen it: SQLite stores float32 as float64 and bool as an integer, MySQL and MariaDB bool as int8. to_dataset now narrows a data variable back to the template's dtype when every value survives exactly, so a plain SELECT round-trips its types while a derived value (an AVG of float32, a SUM of bool) keeps the result's. Only in-memory values are checked; lazily reconstructed variables keep the result's dtype. SQLite also stores times as text, which the chunked round-trip's window predicates cannot compare. The spill now parses text datetime coordinates into timestamps, as it already does for text and interval durations. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- xarray_sql/ds.py | 34 ++++++++++++++++++++++++++++-- xarray_sql/roundtrip.py | 46 +++++++++++++++++++++++------------------ 2 files changed, 58 insertions(+), 22 deletions(-) diff --git a/xarray_sql/ds.py b/xarray_sql/ds.py index 0ad1d958..3803c0ef 100644 --- a/xarray_sql/ds.py +++ b/xarray_sql/ds.py @@ -117,6 +117,33 @@ def _as_dtype(coord: xr.DataArray, dtype: np.dtype) -> xr.DataArray: return coord.astype(dtype) +_NARROWINGS = {("i", "b"), ("u", "b"), ("f", "f")} +"""(result kind, template kind) pairs a database may have widened.""" + + +def _restore_dtype(var: xr.DataArray, dtype: np.dtype) -> xr.DataArray: + """*var* as the template's *dtype*, when that loses nothing. + + Databases without a type widen it: SQLite stores ``float32`` as + ``float64`` and ``bool`` as an integer, MySQL ``bool`` as ``int8``. + A plain ``SELECT`` of such a column narrows back exactly; a derived + value (an ``AVG`` of ``float32``, a ``SUM`` of ``bool``) does not, and + keeps the result's dtype. Only in-memory values can be checked, so a + lazily reconstructed variable keeps the result's dtype too. + """ + if var.dtype == dtype or not isinstance(var.data, np.ndarray): + return var + if (var.dtype.kind, dtype.kind) not in _NARROWINGS: + return var + values = var.values + narrowed = values.astype(dtype) + if not np.array_equal( + narrowed.astype(values.dtype), values, equal_nan=True + ): + return var + return var.copy(data=narrowed) + + def _apply_template(ds: xr.Dataset, template: xr.Dataset) -> xr.Dataset: """Recover metadata that the forward SQL pivot strips. @@ -130,6 +157,8 @@ def _apply_template(ds: xr.Dataset, template: xr.Dataset) -> xr.Dataset: ``AVG`` or a null-introducing filter), and reattaching the source's packing would make a later ``ds.to_netcdf()`` write corrupt values. + * Data-variable dtype, where the database widened it and every value + survives the narrowing exactly (see ``_restore_dtype``). * Dim-coordinate dtype, where SQL upcasted (datetime is the canonical case). * Non-dim coordinates whose dims are all present in ``ds`` (scalar @@ -140,10 +169,11 @@ def _apply_template(ds: xr.Dataset, template: xr.Dataset) -> xr.Dataset: """ out = ds.copy() - # 1. Data-var attrs / encoding for vars present in the template. - # Aggregation aliases absent from template intentionally inherit nothing. + # 1. Data-var dtype, attrs, and encoding for vars present in the + # template. Aggregation aliases absent from template inherit nothing. for name in list(out.data_vars): if name in template.data_vars: + out[name] = _restore_dtype(out[name], template[name].dtype) out[name].attrs = dict(template[name].attrs) # Drop dtype-bound encoding keys; SQL may have changed dtype. enc = { diff --git a/xarray_sql/roundtrip.py b/xarray_sql/roundtrip.py index 9e9d93ef..55b3402e 100644 --- a/xarray_sql/roundtrip.py +++ b/xarray_sql/roundtrip.py @@ -511,15 +511,21 @@ def _as_durations(array: pa.Array) -> pa.Array: return pa.array(pd.to_timedelta(array.to_pandas()), pa.duration("ns")) +def _as_timestamps(array: pa.Array) -> pa.Array: + """Times a database returned as text (SQLite has no time type).""" + return pa.array(pd.to_datetime(array.to_pandas()), pa.timestamp("ns")) + + def _stream_to_parquet( result: Any, path: str, template: xr.Dataset | None = None ) -> None: """Write a one-shot Arrow result to Parquet, batch by batch. - Columns the template holds as ``timedelta64`` coordinates but the - database returned as intervals or text are written as durations: - Parquet cannot store month-day-nano intervals, and the windows the - chunked reconstruction builds compare durations, not text. + Coordinates the template holds as ``timedelta64`` or ``datetime64`` + but the database returned as intervals or text are written as + durations and timestamps: Parquet cannot store month-day-nano + intervals, and the windows the chunked reconstruction builds compare + times, not text. """ opened = _open_stream(result) if opened is None: @@ -528,25 +534,25 @@ def _stream_to_parquet( "Arrow stream." ) schema, batches = opened - durations = [ - i - for i, field in enumerate(schema) - if template is not None - and field.name in template.coords - and template.coords[field.name].dtype.kind == "m" - and ( - pa.types.is_interval(field.type) - or pa.types.is_string(field.type) - or pa.types.is_large_string(field.type) + conversions = {} + for i, field in enumerate(schema): + if template is None or field.name not in template.coords: + continue + kind = template.coords[field.name].dtype.kind + text = pa.types.is_string(field.type) or pa.types.is_large_string( + field.type ) - ] - for i in durations: - schema = schema.set(i, schema.field(i).with_type(pa.duration("ns"))) + if kind == "m" and (text or pa.types.is_interval(field.type)): + conversions[i] = (_as_durations, pa.duration("ns")) + elif kind == "M" and text: + conversions[i] = (_as_timestamps, pa.timestamp("ns")) + for i, (_, arrow_type) in conversions.items(): + schema = schema.set(i, schema.field(i).with_type(arrow_type)) with pq.ParquetWriter(path, schema) as writer: for batch in batches: - if durations: + if conversions: columns = list(batch.columns) - for i in durations: - columns[i] = _as_durations(columns[i]) + for i, (convert, _) in conversions.items(): + columns[i] = convert(columns[i]) batch = pa.RecordBatch.from_arrays(columns, schema=schema) writer.write_batch(batch) From bbc8e327505f28f4a3925efcf28314b99218d404 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 18:18:39 -0700 Subject: [PATCH 11/31] Keep each database's ADBC differences in one dialect table The adapter's per-database handling had grown into scattered vendor checks (ClickHouse table creation, MySQL quoting and default-database switching, integer durations on SQLite and ClickHouse). A Dialect now records those facts for one database, and the adapter has one code path over them: quote character, whether schemas exist and which object holds them, how an ingest reaches a schema, temporary-table and duration support, and table DDL for append-only drivers. A database missing from the table gets standard SQL. Entries for SQLite, DuckDB, PostgreSQL, MySQL/MariaDB, ClickHouse/chDB, DataFusion, and Trino are exercised by the test suite; Spark, BigQuery, Databricks, and Snowflake follow their drivers' published feature tables. Two behavior changes found while testing across drivers: - temporary=True raises ValueError where the driver cannot honor it. Trino's driver silently created a permanent table. - ingest_options= passes driver-specific statement options to each ingest (Spark requires a staging area this way). Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- xarray_sql/backends/_adbc_dialects.py | 242 +++++++++++++++++++++++ xarray_sql/backends/adbc.py | 270 +++++++------------------- 2 files changed, 311 insertions(+), 201 deletions(-) create mode 100644 xarray_sql/backends/_adbc_dialects.py diff --git a/xarray_sql/backends/_adbc_dialects.py b/xarray_sql/backends/_adbc_dialects.py new file mode 100644 index 00000000..8e287291 --- /dev/null +++ b/xarray_sql/backends/_adbc_dialects.py @@ -0,0 +1,242 @@ +"""How the databases behind ADBC drivers differ, in one table. + +ADBC gives every database the same bulk-ingest call, but the SQL around +it and the drivers' support for it vary: how identifiers are quoted, +whether the database has schemas, whether its driver can create tables +or honor ``temporary=True``, and which Arrow types it can store. Each +[Dialect][xarray_sql.backends._adbc_dialects.Dialect] records those +facts for one database, so the adapter has one code path and adding a +database means adding a row. + +A database missing from the table gets the defaults, which are what +ANSI SQL and the ADBC specification prescribe. +""" + +from __future__ import annotations + +import dataclasses +from collections.abc import Callable +from typing import TYPE_CHECKING, Literal + +import pyarrow as pa + +if TYPE_CHECKING: + from adbc_driver_manager import dbapi + +IngestMode = Literal["create", "append", "replace", "create_append"] + +TableDDL = Callable[..., list[str]] +"""Builds the statements that create a table before an append-only ingest.""" + + +@dataclasses.dataclass(frozen=True) +class Dialect: + """What the ADBC adapter needs to know about one database.""" + + name: str + """The database's name, as used in error messages.""" + + quote: str = '"' + """The identifier quote: ``"`` in ANSI SQL, a backtick in MySQL-family + and Hive-family SQL, where ``"era5"`` is a string.""" + + schemas: bool = True + """Whether a mixed-dimension Dataset can go in a schema of its own + (``era5.surface``); otherwise its tables are flat (``era5_surface``).""" + + schema_kind: str = "SCHEMA" + """The object ``CREATE ... IF NOT EXISTS`` makes to hold them.""" + + target_schema: Literal["option", "default_database"] = "option" + """How an ingest reaches a table in that schema: ADBC's target-schema + option, or, where the driver ignores that option, by making the + schema the connection's default database for the ingest.""" + + temporary_tables: bool = True + """Whether the driver honors ``temporary=True``.""" + + durations: bool = True + """Whether the database stores Arrow durations; if not, timedeltas are + ingested as integer counts of their unit.""" + + table_ddl: TableDDL | None = None + """For drivers that can only append: creates each table beforehand.""" + + def quote_identifier(self, identifier: str) -> str: + """*identifier* as a quoted identifier in this database's SQL.""" + escaped = identifier.replace(self.quote, self.quote * 2) + return f"{self.quote}{escaped}{self.quote}" + + +_CLICKHOUSE_TYPES = { + pa.bool_(): "Bool", + pa.int8(): "Int8", + pa.int16(): "Int16", + pa.int32(): "Int32", + pa.int64(): "Int64", + pa.uint8(): "UInt8", + pa.uint16(): "UInt16", + pa.uint32(): "UInt32", + pa.uint64(): "UInt64", + pa.float16(): "Float32", + pa.float32(): "Float32", + pa.float64(): "Float64", + pa.string(): "String", + pa.large_string(): "String", + pa.binary(): "String", + pa.large_binary(): "String", + pa.date32(): "Date32", +} + +_TIMESTAMP_PRECISION = {"s": 0, "ms": 3, "us": 6, "ns": 9} + + +def _clickhouse_type(field: pa.Field, key: bool) -> str: + """The ClickHouse column type for an Arrow field. + + Timestamps without a zone are declared UTC, which is what their + values mean; ClickHouse also parses string literals compared with a + column in that column's zone, so ``time >= '2020-01-01'`` means UTC + rather than the server's local time. Sort-key columns are not + ``Nullable``. Floats are: the scan writes NaN as null so aggregates + skip missing values, and a plain ``Float64`` column would store that + null as 0. + """ + arrow_type = field.type + if pa.types.is_timestamp(arrow_type): + precision = _TIMESTAMP_PRECISION[arrow_type.unit] + zone = arrow_type.tz or "UTC" + name = f"DateTime64({precision}, '{zone}')" + elif arrow_type in _CLICKHOUSE_TYPES: + name = _CLICKHOUSE_TYPES[arrow_type] + else: + raise TypeError( + f"no ClickHouse column type for {field.name!r} of Arrow type " + f"{arrow_type}; create the table yourself and register with " + f'mode="append"' + ) + if key or not field.nullable: + return name + return f"Nullable({name})" + + +def _clickhouse_ddl( + table: str, + schema: pa.Schema, + dims: tuple[str, ...], + *, + mode: IngestMode, + temporary: bool, + database: str | None, +) -> list[str]: + """Statements that prepare *table* for an append-mode ingest. + + ClickHouse's ADBC driver only appends, so the table is created here + for every other mode. Tables are sorted by their dimensions, so + ClickHouse's primary index skips data on dimension predicates the + way chunk pruning does in the other engines. + """ + if mode == "append": + return [] + quote = CLICKHOUSE.quote_identifier + target = quote(table) + if database is not None: + target = f"{quote(database)}.{target}" + columns = ", ".join( + f"{quote(field.name)} {_clickhouse_type(field, field.name in dims)}" + for field in schema + ) + kind = "TEMPORARY TABLE" if temporary else "TABLE" + if temporary: + engine = "ENGINE = Memory" + else: + order = ", ".join(quote(dim) for dim in dims) or "tuple()" + engine = f"ENGINE = MergeTree ORDER BY ({order})" + statements = [] + if mode == "replace": + statements.append(f"DROP {kind} IF EXISTS {target}") + exists = " IF NOT EXISTS" if mode == "create_append" else "" + statements.append(f"CREATE {kind}{exists} {target} ({columns}) {engine}") + return statements + + +CLICKHOUSE = Dialect( + "ClickHouse", + schema_kind="DATABASE", + durations=False, + table_ddl=_clickhouse_ddl, +) + +DIALECTS: dict[str, Dialect] = { + # Exercised by the test suite, against a live database or driver. + "sqlite": Dialect("SQLite", schemas=False, durations=False), + "duckdb": Dialect("DuckDB"), + "postgresql": Dialect("PostgreSQL"), + # MariaDB's server reports itself as MySQL. + "mysql": Dialect( + "MySQL", + quote="`", + schema_kind="DATABASE", + target_schema="default_database", + ), + "clickhouse": CLICKHOUSE, + "datafusion": Dialect("DataFusion", temporary_tables=False), + # Trino's driver ignores temporary=True and creates a permanent table. + "trino": Dialect("Trino", temporary_tables=False), + # From the drivers' published feature tables and the databases' SQL + # references; not exercised by the test suite. + "spark": Dialect("Spark", quote="`", schemas=False, temporary_tables=False), + "bigquery": Dialect("BigQuery", quote="`", temporary_tables=False), + "databricks": Dialect("Databricks", quote="`", temporary_tables=False), + "snowflake": Dialect("Snowflake", temporary_tables=False), +} +"""Known databases, keyed by a name their drivers' vendor names contain.""" + + +def dialect_for(con: dbapi.Connection) -> Dialect: + """The [Dialect][xarray_sql.backends._adbc_dialects.Dialect] of *con*. + + Drivers name their database in ``adbc_get_info``. The ClickHouse + driver does not implement it, so only a driver without it is probed + with a query against ClickHouse's ``system.one`` table. + """ + try: + info = con.adbc_get_info() + except Exception: # noqa: BLE001 — unimplemented; probe instead + pass + else: + vendor = str(info.get("vendor_name") or "").lower() + for key, dialect in DIALECTS.items(): + if key in vendor: + return dialect + return Dialect(str(info.get("vendor_name") or "the database")) + try: + with con.cursor() as cur: + cur.execute("SELECT 1 FROM system.one") + cur.fetchall() + except Exception: # noqa: BLE001 + return Dialect("the database") + return CLICKHOUSE + + +def durations_as_integers( + reader: pa.RecordBatchReader, +) -> pa.RecordBatchReader: + """*reader* with duration columns as integer counts of their unit. + + The template's ``timedelta64`` unit matches the Arrow unit the scan + derived from it, so [xarray_sql.to_dataset][] reads the counts back + as the original durations. + """ + fields = [ + pa.field(f.name, pa.int64(), f.nullable, f.metadata) + if pa.types.is_duration(f.type) + else f + for f in reader.schema + ] + if all(f.type == g.type for f, g in zip(fields, reader.schema)): + return reader + schema = pa.schema(fields, metadata=reader.schema.metadata) + return pa.RecordBatchReader.from_batches( + schema, (batch.cast(schema) for batch in reader) + ) diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index df26e5db..439105a3 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -30,10 +30,9 @@ import contextlib import warnings -from collections.abc import Iterator -from typing import TYPE_CHECKING, Any, Literal, TypeGuard +from collections.abc import Iterator, Mapping +from typing import TYPE_CHECKING, Any, TypeGuard -import pyarrow as pa import xarray as xr from ..df import ( @@ -43,6 +42,12 @@ resolve_table_names, shared_coord_arrays, ) +from ._adbc_dialects import ( + Dialect, + IngestMode, + dialect_for, + durations_as_integers, +) from .base import register_adapter from .pyarrow import XarrayPushdownDataset @@ -51,165 +56,6 @@ __all__ = ["ADBCAdapter"] -IngestMode = Literal["create", "append", "replace", "create_append"] - - -_BACKTICK_VENDORS = ("mysql", "mariadb", "bigquery") -"""Databases whose SQL quotes identifiers with backticks, not ``"``. - -MySQL reads ``"era5"`` as a string unless ``ANSI_QUOTES`` is set, and -BigQuery always does. -""" - - -def _quote(identifier: str, vendor: str = "") -> str: - """Render *identifier* as a quoted SQL identifier in *vendor*'s SQL.""" - if any(name in vendor for name in _BACKTICK_VENDORS): - escaped = identifier.replace("`", "``") - return f"`{escaped}`" - escaped = identifier.replace('"', '""') - return f'"{escaped}"' - - -def _vendor(con: dbapi.Connection) -> str: - """The lowercase name of the database behind *con*; ``""`` if unknown. - - Drivers name their database in ``adbc_get_info``. The ClickHouse - driver does not implement it, so only a driver without it is probed - with a query against ClickHouse's ``system.one`` table. - """ - try: - info = con.adbc_get_info() - except Exception: # noqa: BLE001 — unimplemented; probe instead - pass - else: - return str(info.get("vendor_name") or "").lower() - try: - with con.cursor() as cur: - cur.execute("SELECT 1 FROM system.one") - cur.fetchall() - except Exception: # noqa: BLE001 - return "" - return "clickhouse" - - -_CLICKHOUSE_TYPES = { - pa.bool_(): "Bool", - pa.int8(): "Int8", - pa.int16(): "Int16", - pa.int32(): "Int32", - pa.int64(): "Int64", - pa.uint8(): "UInt8", - pa.uint16(): "UInt16", - pa.uint32(): "UInt32", - pa.uint64(): "UInt64", - pa.float16(): "Float32", - pa.float32(): "Float32", - pa.float64(): "Float64", - pa.string(): "String", - pa.large_string(): "String", - pa.binary(): "String", - pa.large_binary(): "String", - pa.date32(): "Date32", -} - -_TIMESTAMP_PRECISION = {"s": 0, "ms": 3, "us": 6, "ns": 9} - - -def _clickhouse_type(field: pa.Field, key: bool) -> str: - """The ClickHouse column type for an Arrow field. - - Timestamps without a zone are declared UTC, which is what their - values mean; ClickHouse also parses string literals compared with a - column in that column's zone, so ``time >= '2020-01-01'`` means UTC - rather than the server's local time. Sort-key columns are not - ``Nullable``. Floats are: the scan writes NaN as null so aggregates - skip missing values, and a plain ``Float64`` column would store that - null as 0. - """ - arrow_type = field.type - if pa.types.is_timestamp(arrow_type): - precision = _TIMESTAMP_PRECISION[arrow_type.unit] - zone = arrow_type.tz or "UTC" - name = f"DateTime64({precision}, '{zone}')" - elif arrow_type in _CLICKHOUSE_TYPES: - name = _CLICKHOUSE_TYPES[arrow_type] - else: - raise TypeError( - f"no ClickHouse column type for {field.name!r} of Arrow type " - f"{arrow_type}; create the table yourself and register with " - f'mode="append"' - ) - if key or not field.nullable: - return name - return f"Nullable({name})" - - -def _clickhouse_ddl( - table: str, - schema: pa.Schema, - dims: tuple[str, ...], - *, - mode: IngestMode, - temporary: bool, - database: str | None, -) -> list[str]: - """Statements that prepare *table* for an append-mode ingest. - - ClickHouse's ADBC driver only appends, so the table is created here - for every other mode. Tables are sorted by their dimensions, so - ClickHouse's primary index skips data on dimension predicates the - way chunk pruning does in the other engines. - """ - if mode == "append": - return [] - target = _quote(table) - if database is not None: - target = f"{_quote(database)}.{target}" - columns = ", ".join( - f"{_quote(field.name)} {_clickhouse_type(field, field.name in dims)}" - for field in schema - ) - kind = "TEMPORARY TABLE" if temporary else "TABLE" - if temporary: - engine = "ENGINE = Memory" - else: - order = ", ".join(_quote(dim) for dim in dims) or "tuple()" - engine = f"ENGINE = MergeTree ORDER BY ({order})" - statements = [] - if mode == "replace": - statements.append(f"DROP {kind} IF EXISTS {target}") - exists = " IF NOT EXISTS" if mode == "create_append" else "" - statements.append(f"CREATE {kind}{exists} {target} ({columns}) {engine}") - return statements - - -_NO_DURATION_VENDORS = ("sqlite", "clickhouse") -"""Databases with no duration type; timedeltas are stored as integers.""" - - -def _durations_as_integers( - reader: pa.RecordBatchReader, -) -> pa.RecordBatchReader: - """*reader* with duration columns as integer counts of their unit. - - The template's ``timedelta64`` unit matches the Arrow unit the scan - derived from it, so [xarray_sql.to_dataset][] reads the counts back - as the original durations. - """ - fields = [ - pa.field(f.name, pa.int64(), f.nullable, f.metadata) - if pa.types.is_duration(f.type) - else f - for f in reader.schema - ] - if all(f.type == g.type for f, g in zip(fields, reader.schema)): - return reader - schema = pa.schema(fields, metadata=reader.schema.metadata) - return pa.RecordBatchReader.from_batches( - schema, (batch.cast(schema) for batch in reader) - ) - def _ingest( con: dbapi.Connection, @@ -217,10 +63,11 @@ def _ingest( ds: xr.Dataset, chunks: Chunks, *, + dialect: Dialect, dims: tuple[str, ...], mode: IngestMode, temporary: bool, - vendor: str, + ingest_options: Mapping[str, str] | None, db_schema_name: str | None = None, **kwargs: Any, ) -> None: @@ -232,10 +79,10 @@ def _ingest( batches, so the source read and the database write overlap. """ reader = XarrayPushdownDataset(ds, chunks, **kwargs).scanner().to_reader() - if vendor in _NO_DURATION_VENDORS: - reader = _durations_as_integers(reader) - if vendor == "clickhouse": - for statement in _clickhouse_ddl( + if not dialect.durations: + reader = durations_as_integers(reader) + if dialect.table_ddl is not None: + for statement in dialect.table_ddl( table, reader.schema, dims, @@ -250,13 +97,15 @@ def _ingest( in_database: contextlib.AbstractContextManager[None] = ( contextlib.nullcontext() ) - if db_schema_name is not None and vendor.startswith(("mysql", "mariadb")): - # The MySQL driver creates the table in the connection's default - # database but inserts into the one it was given, so name the - # target by making it the default instead. - in_database = _default_database(con, db_schema_name, vendor) + if ( + db_schema_name is not None + and dialect.target_schema == "default_database" + ): + in_database = _default_database(con, db_schema_name, dialect) target_schema = None with in_database, con.cursor() as cur: + if ingest_options: + cur.adbc_statement.set_options(**ingest_options) cur.adbc_ingest( table, reader, @@ -268,23 +117,25 @@ def _ingest( @contextlib.contextmanager def _default_database( - con: dbapi.Connection, database: str, vendor: str + con: dbapi.Connection, database: str, dialect: Dialect ) -> Iterator[None]: - """Make *database* the MySQL connection's default while in the block. + """Make *database* the connection's default while in the block. - The previous default is restored afterwards. MySQL cannot unset a - default database, so a connection that had none keeps *database*. + For drivers that create an ingest's table in the default database + whatever target schema they are given. The previous default is + restored afterwards; MySQL cannot unset a default database, so a + connection that had none keeps *database*. """ with con.cursor() as cur: cur.execute("SELECT DATABASE()") (previous,) = cur.fetchone() - cur.execute(f"USE {_quote(database, vendor)}") + cur.execute(f"USE {dialect.quote_identifier(database)}") try: yield finally: if previous is not None: with con.cursor() as cur: - cur.execute(f"USE {_quote(previous, vendor)}") + cur.execute(f"USE {dialect.quote_identifier(previous)}") def _schema_exists(con: dbapi.Connection, name: str) -> bool: @@ -313,27 +164,37 @@ def _connection_usable(con: dbapi.Connection) -> bool: return True -def _create_schema(con: dbapi.Connection, name: str, *, vendor: str) -> bool: +def _flat(name: str, reason: str) -> bool: + """Warn that *name*'s groups become flat tables; always ``False``.""" + warnings.warn( + f"Registering the dimension groups of {name!r} as flat " + f"{name}_ tables: {reason}.", + RuntimeWarning, + stacklevel=4, + ) + return False + + +def _create_schema(con: dbapi.Connection, name: str, dialect: Dialect) -> bool: """Ensure the database schema *name* exists; whether it does. An existing schema is used as is: creating it can need privileges on the whole database (PostgreSQL checks them even for ``IF NOT EXISTS``) that a role granted only that schema lacks. - Not every ADBC database has schemas (SQLite does not), and creating - one can fail for lack of privileges. When the connection survives the - failure, the caller falls back to flat table names. On databases - where a failed statement aborts the transaction (PostgreSQL), nothing - after it could run, so this raises instead of falling back. In - ClickHouse, a database plays the role of a schema. + Creating one can fail for lack of privileges. When the connection + survives the failure, the caller falls back to flat table names. On + databases where a failed statement aborts the transaction + (PostgreSQL), nothing after it could run, so this raises instead. """ + if not dialect.schemas: + return _flat(name, f"{dialect.name} has no schemas to hold them") if _schema_exists(con, name): return True - kind = "DATABASE" if vendor == "clickhouse" else "SCHEMA" - target = _quote(name, vendor) + target = dialect.quote_identifier(name) try: with con.cursor() as cur: - cur.execute(f"CREATE {kind} IF NOT EXISTS {target}") + cur.execute(f"CREATE {dialect.schema_kind} IF NOT EXISTS {target}") except Exception as exc: if not _connection_usable(con): raise RuntimeError( @@ -344,14 +205,7 @@ def _create_schema(con: dbapi.Connection, name: str, *, vendor: str) -> bool: f"the privilege to), or pass temporary=True to register " f"flat {name}_ tables instead." ) from exc - warnings.warn( - f"Could not create the {name!r} schema to hold the dimension " - f"groups of {name!r} ({exc}); registering them as flat " - f"{name}_ tables instead.", - RuntimeWarning, - stacklevel=4, - ) - return False + return _flat(name, f"could not create the {name!r} schema ({exc})") return True @@ -379,6 +233,7 @@ def register( table_names: TableNames = None, mode: IngestMode = "create", temporary: bool = False, + ingest_options: Mapping[str, str] | None = None, **kwargs: Any, ) -> dbapi.Connection: """Ingest ``ds`` into tables on an ADBC connection. @@ -421,7 +276,12 @@ def register( ``"replace"`` drops and recreates it, ``"append"`` and ``"create_append"`` add rows to it. temporary: Create temporary tables, which the database drops - when the connection closes. + when the connection closes. Raises ``ValueError`` on + databases whose driver cannot (DataFusion, Trino, Spark, + BigQuery, Databricks, Snowflake). + ingest_options: Driver-specific statement options set on each + ingest, e.g. Spark's + ``{"spark.ingest.staging_area_uri": "s3://bucket/path"}``. **kwargs: Forwarded to [XarrayPushdownDataset][xarray_sql.backends.pyarrow.XarrayPushdownDataset] (``batch_size``, ``prefetch``, ``prefetch_bytes``, @@ -429,22 +289,29 @@ def register( """ groups = group_vars_by_dims(ds) names = resolve_table_names(ds, table_names, case_insensitive=True) - vendor = _vendor(con) + dialect = dialect_for(con) + if temporary and not dialect.temporary_tables: + raise ValueError( + f"{dialect.name}'s ADBC driver does not support temporary " + f"tables; register without temporary=True and drop the " + f"tables when done." + ) if len(groups) <= 1: _ingest( con, name, ds, chunks, + dialect=dialect, dims=next(iter(groups), ()), mode=mode, temporary=temporary, - vendor=vendor, + ingest_options=ingest_options, **kwargs, ) return con - in_schema = not temporary and _create_schema(con, name, vendor=vendor) + in_schema = not temporary and _create_schema(con, name, dialect) coord_arrays = shared_coord_arrays(ds) for dims, var_names in groups.items(): group = names[dims] @@ -453,10 +320,11 @@ def register( group if in_schema else f"{name}_{group}", ds[var_names], chunks, + dialect=dialect, dims=dims, mode=mode, temporary=temporary, - vendor=vendor, + ingest_options=ingest_options, db_schema_name=name if in_schema else None, coord_arrays=coord_arrays, **kwargs, From 6738827b9aef39f6151655708afa1bc90e133423 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 18:18:39 -0700 Subject: [PATCH 12/31] Run one ADBC contract suite against every available database The ADBC tests were per-database copies with fixtures spread through the module. They are now one contract, run once per backend in a BACKENDS table at the top of the module: round-trips with exact dtypes and missing values, every mode, temporary tables, mixed-dimension naming, timedelta coordinates, and the chunked round-trip. Tests of a single database's specifics follow. SQLite and DuckDB always run; chDB and DataFusion run in-process when their drivers are installed (CI now installs both); ClickHouse, PostgreSQL, MySQL, MariaDB, and Trino run when XARRAY_SQL_TEST__URI points at a server (PostgreSQL's variable is now XARRAY_SQL_TEST_POSTGRESQL_URI, matching the others). The unit test of the ClickHouse DDL is dropped: the public missing-value test it guarded now runs against chDB in CI. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- .github/workflows/ci.yml | 10 +- tests/test_adbc_backend.py | 622 ++++++++++++++++++------------------- 2 files changed, 315 insertions(+), 317 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index aafb25d8..eb57c0f4 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -68,17 +68,17 @@ jobs: run: uv sync --dev --no-install-package xarray-sql - name: build rust run: uv run --no-project maturin develop --uv - - name: Install the chDB ADBC driver - # Embedded ClickHouse: runs the ClickHouse adapter tests without a - # server. dbc installs into the active virtualenv. + - name: Install in-process ADBC drivers + # chDB (embedded ClickHouse) and DataFusion run the ADBC adapter's + # contract tests without a server. dbc installs into the active + # virtualenv. env: VIRTUAL_ENV: ${{ github.workspace }}/.venv run: | uv pip install dbc uv run --no-project dbc install chdb + uv run --no-project dbc install datafusion - name: Run unit tests env: VIRTUAL_ENV: ${{ github.workspace }}/.venv - XARRAY_SQL_TEST_CLICKHOUSE_DRIVER: chdb - XARRAY_SQL_TEST_CLICKHOUSE_URI: "chdb://" run: uv run --no-project pytest -v . -m "not integration" diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index ba9bbe42..ac52e8d2 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -1,35 +1,45 @@ -"""Tests for the ADBC engine adapter. +"""Tests for the ADBC engine adapter, run against every available database. -``xql.register`` ingests a Dataset into any database reachable through an -ADBC driver, and ``xql.to_dataset`` rebuilds a labeled Dataset from the -driver's Arrow cursor. Two drivers cover the two shapes of database: -SQLite, which has no schemas and stores timestamps as text, and DuckDB's -built-in ADBC driver, which has schemas and native temporal types. +``xql.register`` ingests a Dataset into any database with an ADBC driver, +and ``xql.to_dataset`` rebuilds a labeled Dataset from the driver's Arrow +cursor. The contract tests below run once per backend in ``BACKENDS``; +tests of one database's specifics follow them. + +SQLite and DuckDB always run. The others run when available: + +- ``chdb`` and ``datafusion`` run in-process once their drivers are + installed (``dbc install chdb datafusion``). +- ``clickhouse``, ``postgresql``, ``mysql``, ``mariadb``, and ``trino`` + need a server: set ``XARRAY_SQL_TEST__URI`` to its URI (and + ``XARRAY_SQL_TEST_CLICKHOUSE_DRIVER`` for a ClickHouse driver that + ``dbc`` did not install). """ +import dataclasses import importlib.util import os +import uuid import numpy as np import pandas as pd -import pyarrow as pa import pytest import xarray as xr import xarray_sql as xql -from xarray_sql.backends.adbc import _clickhouse_ddl dbapi = pytest.importorskip("adbc_driver_manager.dbapi") -sqlite_dbapi = pytest.importorskip("adbc_driver_sqlite.dbapi") -NAMES = { - ("time", "lat", "lon"): "surface", - ("time", "level", "lat", "lon"): "atmosphere", -} + +def _module_driver(module: str) -> str | None: + """The driver library an ``adbc_driver_*`` Python package ships.""" + try: + return importlib.import_module(module)._driver_path() + except ImportError: + return None -def _duckdb_driver_path() -> str | None: - """Path of the shared library holding DuckDB's ADBC entrypoint.""" +def _duckdb_driver() -> str | None: + """The shared library holding DuckDB's ADBC entrypoint.""" for module in ("_duckdb", "duckdb.duckdb"): try: spec = importlib.util.find_spec(module) @@ -40,43 +50,156 @@ def _duckdb_driver_path() -> str | None: return None -def _connect(driver: str): - if driver == "sqlite": - return sqlite_dbapi.connect() - path = _duckdb_driver_path() - if path is None: - pytest.skip("duckdb is not installed") - return dbapi.connect(driver=path, entrypoint="duckdb_adbc_init") - - -@pytest.fixture(params=["sqlite", "duckdb"]) -def con(request): - connection = _connect(request.param) - yield connection - connection.close() - +def _env(name: str) -> str | None: + return os.environ.get(f"XARRAY_SQL_TEST_{name.upper()}_URI") + + +@dataclasses.dataclass(frozen=True) +class Backend: + """A database to run the contract tests against.""" + + name: str + driver: str | None + uri: str | None = None + entrypoint: str | None = None + schemas: bool = True + """Whether mixed-dimension Datasets register as ``name.group``.""" + temporary: bool = True + """Whether ``temporary=True`` is supported.""" + drop_schema: str = "DROP SCHEMA IF EXISTS {} CASCADE" + quote: str = '"' + needs_uri: bool = False + + def connect(self): + if self.driver is None or (self.needs_uri and not self.uri): + pytest.skip(f"{self.name} is not available; see module docstring") + kwargs = {"db_kwargs": {"uri": self.uri}} if self.uri else {} + if self.entrypoint: + kwargs["entrypoint"] = self.entrypoint + try: + return dbapi.connect(driver=self.driver, **kwargs) + except dbapi.Error as exc: + if self.needs_uri: + raise + pytest.skip(f"{self.name} driver is not installed ({exc})") + + +BACKENDS = [ + Backend("sqlite", _module_driver("adbc_driver_sqlite"), schemas=False), + Backend("duckdb", _duckdb_driver(), entrypoint="duckdb_adbc_init"), + Backend( + "chdb", + "chdb", + uri="chdb://", + drop_schema="DROP DATABASE IF EXISTS {}", + ), + Backend("datafusion", "datafusion", temporary=False), + Backend( + "clickhouse", + os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_DRIVER", "clickhouse"), + uri=_env("clickhouse"), + drop_schema="DROP DATABASE IF EXISTS {}", + needs_uri=True, + ), + Backend( + "postgresql", + _module_driver("adbc_driver_postgresql") or "postgresql", + uri=_env("postgresql"), + needs_uri=True, + ), + Backend( + "mysql", + "mysql", + uri=_env("mysql"), + drop_schema="DROP DATABASE IF EXISTS {}", + quote="`", + needs_uri=True, + ), + Backend( + "mariadb", + "mysql", + uri=_env("mariadb"), + drop_schema="DROP DATABASE IF EXISTS {}", + quote="`", + needs_uri=True, + ), + Backend( + "trino", + "trino", + uri=_env("trino"), + temporary=False, + needs_uri=True, + ), +] -@pytest.fixture -def sqlite_con(): - connection = _connect("sqlite") - yield connection - connection.close() +NAMES = { + ("time", "lat", "lon"): "surface", + ("time", "level", "lat", "lon"): "atmosphere", +} -@pytest.fixture -def duckdb_con(): - connection = _connect("duckdb") - yield connection - connection.close() +class Database: + """A connection plus the unique names a test creates, dropped after.""" + + def __init__(self, backend: Backend, con) -> None: + self.backend = backend + self.con = con + self._created: list[str] = [] + + def name(self, base: str) -> str: + """A fresh table (or schema) name, dropped when the test ends.""" + name = f"{base}_{uuid.uuid4().hex[:8]}" + self._created.append(name) + return name + + def query(self, sql: str): + cur = self.con.cursor() + cur.execute(sql) + return cur + + def cleanup(self) -> None: + postgresql = self.backend.name == "postgresql" + if postgresql: + self.con.rollback() + for name in self._created: + quoted = f"{self.backend.quote}{name}{self.backend.quote}" + for statement in ( + f"DROP TABLE IF EXISTS {quoted}", + self.backend.drop_schema.format(quoted), + ): + try: + self.query(statement).close() + except dbapi.Error: + if postgresql: + self.con.rollback() + if postgresql: + self.con.commit() + + +@pytest.fixture(params=BACKENDS, ids=[b.name for b in BACKENDS]) +def db(request): + backend = request.param + database = Database(backend, backend.connect()) + yield database + database.cleanup() + database.con.close() @pytest.fixture def ds() -> xr.Dataset: - np.random.seed(3) - return xr.Dataset( + rng = np.random.default_rng(3) + weather = xr.Dataset( data_vars=dict( - temperature=(["time", "lat", "lon"], np.random.randn(8, 5, 6)), - precipitation=(["time", "lat", "lon"], np.random.rand(8, 5, 6)), + temperature=( + ["time", "lat", "lon"], + rng.standard_normal((8, 5, 6)), + ), + count=(["time", "lat", "lon"], rng.integers(0, 100, (8, 5, 6))), + sst=( + ["time", "lat", "lon"], + rng.random((8, 5, 6)).astype("float32"), + ), + land=(["time", "lat", "lon"], rng.random((8, 5, 6)) > 0.5), ), coords=dict( time=pd.date_range("2021-01-01", periods=8, freq="h"), @@ -85,17 +208,19 @@ def ds() -> xr.Dataset: ), attrs=dict(description="Synthetic weather."), ).chunk({"time": 4}) + weather["temperature"][0, 0, 0] = np.nan + return weather @pytest.fixture def mixed_ds() -> xr.Dataset: - np.random.seed(11) + rng = np.random.default_rng(11) return xr.Dataset( { - "t2m": (["time", "lat", "lon"], np.random.rand(6, 3, 4)), + "t2m": (["time", "lat", "lon"], rng.random((6, 3, 4))), "temperature": ( ["time", "level", "lat", "lon"], - np.random.rand(6, 2, 3, 4), + rng.random((6, 2, 3, 4)), ), }, coords={ @@ -107,33 +232,46 @@ def mixed_ds() -> xr.Dataset: ).chunk({"time": 2}) -def _query(con, sql: str): - cur = con.cursor() - cur.execute(sql) - return cur - +@pytest.fixture +def forecast() -> xr.Dataset: + rng = np.random.default_rng(7) + return xr.Dataset( + {"t2m": (["step", "lat"], rng.random((4, 2)))}, + coords={ + "step": pd.to_timedelta([0, 6, 12, 18], unit="h"), + "lat": [1.0, 2.0], + }, + ).chunk({"step": 2}) -def test_full_scan_round_trips(con, ds): - xql.register(con, "weather", ds) - cur = _query( - con, - "SELECT time, lat, lon, temperature, precipitation FROM weather " - "ORDER BY time, lat, lon", +def _select_all(db, table: str): + return db.query( + f"SELECT time, lat, lon, temperature, count, sst, land FROM {table} " + "ORDER BY time, lat, lon" ) - out = xql.to_dataset(cur, template=ds) - xr.testing.assert_allclose(out, ds.compute()) - assert out.attrs == ds.attrs +# The contract, on every backend ------------------------------------------- + + +def test_round_trip_keeps_values_and_dtypes(db, ds): + table = db.name("weather") + xql.register(db.con, table, ds) + + out = xql.to_dataset(_select_all(db, table), template=ds) + + # NaN survives as missing, and bool/float32 come back as themselves + # even where the database widens them (SQLite, MySQL). + xr.testing.assert_identical(out, ds.compute()) -def test_aggregation_round_trips_on_surviving_dims(con, ds): - xql.register(con, "weather", ds) - cur = _query( - con, - "SELECT lat, lon, AVG(temperature) AS temperature FROM weather " - "GROUP BY lat, lon ORDER BY lat, lon", +def test_aggregates_skip_missing_values(db, ds): + table = db.name("weather") + xql.register(db.con, table, ds) + + cur = db.query( + f"SELECT lat, lon, AVG(temperature) AS temperature FROM {table} " + "GROUP BY lat, lon ORDER BY lat, lon" ) out = xql.to_dataset(cur, template=ds) @@ -143,207 +281,151 @@ def test_aggregation_round_trips_on_surviving_dims(con, ds): ) -def test_chunked_round_trip_spills_the_cursor(duckdb_con, ds): - # DuckDB keeps `time` a timestamp; SQLite returns it as text, which - # only the eager round-trip recovers. - xql.register(duckdb_con, "weather", ds) +def test_existing_table_is_not_overwritten_by_default(db, ds): + table = db.name("weather") + xql.register(db.con, table, ds) - cur = _query( - duckdb_con, - "SELECT time, lat, lon, temperature FROM weather " - "ORDER BY time, lat, lon", - ) - out = xql.to_dataset(cur, template=ds, chunks={"time": 2}, spill=True) + with pytest.raises(dbapi.Error): + xql.register(db.con, table, ds) - assert out.temperature.chunks is not None - xr.testing.assert_allclose( - out.temperature.compute(), ds.temperature.compute() - ) +def test_replace_then_append(db, ds): + table = db.name("weather") + xql.register(db.con, table, ds) + xql.register(db.con, table, ds.isel(time=slice(0, 4)), mode="replace") + xql.register(db.con, table, ds.isel(time=slice(4, 8)), mode="append") -@pytest.mark.parametrize("chunks", [None, {"step": 2}]) -def test_timedelta_coordinates_round_trip(con, chunks): - # Forecast `step`: SQLite stores it as an integer count, DuckDB - # returns it as an interval. - forecast = xr.Dataset( - {"t2m": (["step", "lat"], np.random.rand(4, 2))}, - coords={ - "step": pd.to_timedelta([0, 6, 12, 18], unit="h"), - "lat": [1.0, 2.0], - }, - ).chunk({"step": 2}) - xql.register(con, "forecast", forecast) + out = xql.to_dataset(_select_all(db, table), template=ds) - cur = _query(con, "SELECT step, lat, t2m FROM forecast ORDER BY step, lat") - out = xql.to_dataset( - cur, template=forecast, chunks=chunks, spill=chunks is not None - ) + xr.testing.assert_identical(out, ds.compute()) - xr.testing.assert_identical(out.compute().step, forecast.step) - xr.testing.assert_allclose(out.compute(), forecast.compute()) +def test_create_append_creates_then_appends(db, ds): + table = db.name("weather") + xql.register(db.con, table, ds, mode="create_append") + xql.register(db.con, table, ds, mode="create_append") -def test_existing_table_is_not_overwritten_by_default(con, ds): - xql.register(con, "weather", ds) + count = db.query(f"SELECT COUNT(*) FROM {table}").fetchone()[0] + assert count == 2 * 8 * 5 * 6 - with pytest.raises(dbapi.Error): - xql.register(con, "weather", ds) +def test_temporary_tables(db, ds): + table = db.name("weather") + if not db.backend.temporary: + with pytest.raises(ValueError, match="temporary"): + xql.register(db.con, table, ds, temporary=True) + return -def test_replace_mode_recreates_the_table(con, ds): - xql.register(con, "weather", ds) - xql.register(con, "weather", ds.isel(time=slice(0, 4)), mode="replace") + xql.register(db.con, table, ds, temporary=True) - count = _query(con, "SELECT COUNT(*) FROM weather").fetchone()[0] - assert count == 4 * 5 * 6 + count = db.query(f"SELECT COUNT(*) FROM {table}").fetchone()[0] + assert count == 8 * 5 * 6 -def test_append_mode_adds_rows(con, ds): - xql.register(con, "weather", ds.isel(time=slice(0, 4))) - xql.register(con, "weather", ds.isel(time=slice(4, 8)), mode="append") +def test_mixed_dimensions_are_named_like_every_engine(db, mixed_ds): + name = db.name("era5") + if db.backend.schemas: + xql.register(db.con, name, mixed_ds, table_names=NAMES) + table = f"{name}.atmosphere" + else: + with pytest.warns(RuntimeWarning, match="flat"): + xql.register(db.con, name, mixed_ds, table_names=NAMES) + table = f"{name}_atmosphere" - cur = _query( - con, - "SELECT time, lat, lon, temperature, precipitation FROM weather " - "ORDER BY time, lat, lon", + cur = db.query( + f"SELECT time, level, lat, lon, temperature FROM {table} " + "ORDER BY time, level, lat, lon" ) - xr.testing.assert_allclose(xql.to_dataset(cur, template=ds), ds.compute()) - - -def test_temporary_table_is_queryable(con, ds): - xql.register(con, "weather", ds, temporary=True) + out = xql.to_dataset(cur, template=mixed_ds[["temperature"]]) - count = _query(con, "SELECT COUNT(*) FROM weather").fetchone()[0] - assert count == 8 * 5 * 6 + xr.testing.assert_allclose(out.temperature, mixed_ds.temperature.compute()) -def test_mixed_dimensions_register_in_a_schema(duckdb_con, mixed_ds): - xql.register(duckdb_con, "era5", mixed_ds, table_names=NAMES) +@pytest.mark.parametrize( + "chunks", [None, {"step": 2}], ids=["eager", "chunked"] +) +def test_timedelta_coordinates_round_trip(db, forecast, chunks): + # Stored as a duration, an integer count (SQLite, ClickHouse), an + # interval (DuckDB, PostgreSQL), or text (MySQL, Trino). + table = db.name("forecast") + xql.register(db.con, table, forecast) - cur = _query(duckdb_con, "SELECT AVG(t2m) FROM era5.surface") - assert cur.fetchone()[0] == pytest.approx(float(mixed_ds.t2m.mean())) - cur = _query( - duckdb_con, - "SELECT time, level, lat, lon, temperature FROM era5.atmosphere " - "ORDER BY time, level, lat, lon", + cur = db.query(f"SELECT step, lat, t2m FROM {table} ORDER BY step, lat") + out = xql.to_dataset( + cur, template=forecast, chunks=chunks, spill=chunks is not None ) - out = xql.to_dataset(cur, template=mixed_ds[["temperature"]]) - xr.testing.assert_allclose(out.temperature, mixed_ds.temperature.compute()) + xr.testing.assert_allclose(out.compute(), forecast.compute()) + assert out.step.dtype == forecast.step.dtype -def test_mixed_dimensions_fall_back_to_flat_names(sqlite_con, mixed_ds): - with pytest.warns(RuntimeWarning, match="flat"): - xql.register(sqlite_con, "era5", mixed_ds, table_names=NAMES) - cur = _query(sqlite_con, "SELECT AVG(t2m) FROM era5_surface") - assert cur.fetchone()[0] == pytest.approx(float(mixed_ds.t2m.mean())) - count = _query(sqlite_con, "SELECT COUNT(*) FROM era5_atmosphere") - assert count.fetchone()[0] == 6 * 2 * 3 * 4 +def test_chunked_round_trip_spills_the_cursor(db, ds): + table = db.name("weather") + xql.register(db.con, table, ds) + cur = db.query( + f"SELECT time, lat, lon, temperature FROM {table} " + "ORDER BY time, lat, lon" + ) + out = xql.to_dataset(cur, template=ds, chunks={"time": 2}, spill=True) -def test_temporary_mixed_dimensions_use_flat_names(duckdb_con, mixed_ds): - xql.register( - duckdb_con, "era5", mixed_ds, table_names=NAMES, temporary=True + assert out.temperature.chunks is not None + xr.testing.assert_allclose( + out.temperature.compute(), ds.temperature.compute() ) - count = _query(duckdb_con, "SELECT COUNT(*) FROM era5_surface") - assert count.fetchone()[0] == 6 * 3 * 4 +# One database's specifics --------------------------------------------------- -def test_clickhouse_float_columns_are_nullable(): - # The scan writes NaN as an Arrow null, so aggregates skip it; a plain - # Float64 column would store that null as 0. - schema = pa.schema( - [ - pa.field("time", pa.timestamp("ns")), - pa.field("lat", pa.float64()), - pa.field("t2m", pa.float64()), - pa.field("sst", pa.float32()), - ] - ) - [ddl] = _clickhouse_ddl( - "weather", - schema, - ("time", "lat"), - mode="create", - temporary=False, - database=None, - ) - assert '"t2m" Nullable(Float64)' in ddl - assert '"sst" Nullable(Float32)' in ddl - assert '"lat" Float64,' in ddl # sort keys stay non-Nullable +def _only(db, *names: str) -> None: + if db.backend.name not in names: + pytest.skip(f"specific to {', '.join(names)}") -@pytest.fixture -def postgres_con(): - uri = os.environ.get("XARRAY_SQL_TEST_POSTGRES_URI") - if not uri: - pytest.skip( - "set XARRAY_SQL_TEST_POSTGRES_URI to run against PostgreSQL" + +def test_ingest_options_reach_the_driver(db, ds): + _only(db, "sqlite") + + with pytest.raises(dbapi.Error, match="not.an.option"): + xql.register( + db.con, + db.name("weather"), + ds, + ingest_options={"not.an.option": "x"}, ) - postgres = pytest.importorskip("adbc_driver_postgresql.dbapi") - connection = postgres.connect(uri) - yield connection - connection.rollback() - connection.close() -def test_postgres_schema_failure_explains_the_aborted_transaction( - postgres_con, mixed_ds +def test_postgresql_schema_failure_explains_the_aborted_transaction( + db, mixed_ds ): # PostgreSQL rejects schema names starting with `pg_`, so CREATE SCHEMA # fails here for any user, and a failed statement aborts the # transaction: no fallback ingest could run after it. - with pytest.raises(RuntimeError, match="rollback"): - xql.register(postgres_con, "pg_era5", mixed_ds, table_names=NAMES) + _only(db, "postgresql") + with pytest.raises(RuntimeError, match="rollback"): + xql.register(db.con, db.name("pg_era5"), mixed_ds, table_names=NAMES) -def test_postgres_uses_an_existing_schema(postgres_con, mixed_ds): - with postgres_con.cursor() as cur: - cur.execute("CREATE SCHEMA IF NOT EXISTS era5") - xql.register(postgres_con, "era5", mixed_ds, table_names=NAMES) - count = _query(postgres_con, "SELECT COUNT(*) FROM era5.surface") - assert count.fetchone()[0] == 6 * 3 * 4 +def test_postgresql_uses_an_existing_schema(db, mixed_ds): + _only(db, "postgresql") + name = db.name("era5") + db.query(f'CREATE SCHEMA "{name}"').close() + xql.register(db.con, name, mixed_ds, table_names=NAMES) -@pytest.fixture -def clickhouse_con(): - uri = os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_URI") - driver = os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_DRIVER") - if not (uri and driver): - pytest.skip( - "set XARRAY_SQL_TEST_CLICKHOUSE_URI and " - "XARRAY_SQL_TEST_CLICKHOUSE_DRIVER to run against ClickHouse" - ) - connection = dbapi.connect(driver=driver, db_kwargs={"uri": uri}) - for statement in [ - "DROP TABLE IF EXISTS weather", - "DROP DATABASE IF EXISTS era5", - ]: - _query(connection, statement).close() - yield connection - connection.close() - - -def test_clickhouse_round_trips(clickhouse_con, ds): - xql.register(clickhouse_con, "weather", ds) - - cur = _query( - clickhouse_con, - "SELECT time, lat, lon, temperature, precipitation FROM weather " - "ORDER BY time, lat, lon", - ) - xr.testing.assert_allclose(xql.to_dataset(cur, template=ds), ds.compute()) + count = db.query(f"SELECT COUNT(*) FROM {name}.surface").fetchone()[0] + assert count == 6 * 3 * 4 -def test_clickhouse_time_literals_mean_utc(clickhouse_con, ds): - xql.register(clickhouse_con, "weather", ds) +def test_clickhouse_time_literals_mean_utc(db, ds): + _only(db, "clickhouse", "chdb") + table = db.name("weather") + xql.register(db.con, table, ds) - cur = _query( - clickhouse_con, - "SELECT time, lat, lon, temperature FROM weather " - "WHERE time >= '2021-01-01 04:00:00' ORDER BY time, lat, lon", + cur = db.query( + f"SELECT time, lat, lon, temperature FROM {table} " + "WHERE time >= '2021-01-01 04:00:00' ORDER BY time, lat, lon" ) out = xql.to_dataset(cur, template=ds) @@ -351,96 +433,12 @@ def test_clickhouse_time_literals_mean_utc(clickhouse_con, ds): xr.testing.assert_allclose(out.temperature, expected.compute()) -def test_clickhouse_replace_then_append(clickhouse_con, ds): - xql.register(clickhouse_con, "weather", ds) - xql.register( - clickhouse_con, "weather", ds.isel(time=slice(0, 4)), mode="replace" - ) - xql.register( - clickhouse_con, "weather", ds.isel(time=slice(4, 8)), mode="append" - ) - - cur = _query( - clickhouse_con, - "SELECT time, lat, lon, temperature, precipitation FROM weather " - "ORDER BY time, lat, lon", - ) - xr.testing.assert_allclose(xql.to_dataset(cur, template=ds), ds.compute()) - - -def test_clickhouse_temporary_table_is_queryable(clickhouse_con, ds): - xql.register(clickhouse_con, "weather", ds, temporary=True) - - count = _query(clickhouse_con, "SELECT COUNT(*) FROM weather") - assert count.fetchone()[0] == 8 * 5 * 6 - - -def test_clickhouse_mixed_dimensions_register_in_a_database( - clickhouse_con, mixed_ds -): - xql.register(clickhouse_con, "era5", mixed_ds, table_names=NAMES) - - cur = _query( - clickhouse_con, - "SELECT time, level, lat, lon, temperature FROM era5.atmosphere " - "ORDER BY time, level, lat, lon", - ) - out = xql.to_dataset(cur, template=mixed_ds[["temperature"]]) - xr.testing.assert_allclose(out.temperature, mixed_ds.temperature.compute()) - +def test_mysql_keeps_the_default_database(db, mixed_ds): + # The MySQL driver ignores the target schema, so the adapter switches + # the default database for the ingest and must switch it back. + _only(db, "mysql", "mariadb") + before = db.query("SELECT DATABASE()").fetchone()[0] -def test_clickhouse_keeps_missing_values(clickhouse_con, ds): - holed = ds.copy(deep=True) - holed["temperature"][0, 0, 0] = np.nan - xql.register(clickhouse_con, "weather", holed) - - with _query(clickhouse_con, "SELECT AVG(temperature) FROM weather") as avg: - mean = avg.fetchone()[0] - assert mean == pytest.approx(float(holed.temperature.mean())) - cur = _query( - clickhouse_con, - "SELECT time, lat, lon, temperature FROM weather " - "ORDER BY time, lat, lon", - ) - out = xql.to_dataset(cur, template=holed) - xr.testing.assert_allclose(out.temperature, holed.temperature.compute()) + xql.register(db.con, db.name("era5"), mixed_ds, table_names=NAMES) - -@pytest.fixture -def mysql_con(): - uri = os.environ.get("XARRAY_SQL_TEST_MYSQL_URI") - if not uri: - pytest.skip("set XARRAY_SQL_TEST_MYSQL_URI to run against MySQL") - connection = dbapi.connect(driver="mysql", db_kwargs={"uri": uri}) - for statement in [ - "DROP TABLE IF EXISTS weather", - "DROP DATABASE IF EXISTS era5", - ]: - _query(connection, statement).close() - yield connection - connection.close() - - -def test_mysql_round_trips(mysql_con, ds): - xql.register(mysql_con, "weather", ds) - - cur = _query( - mysql_con, - "SELECT time, lat, lon, temperature, precipitation FROM weather " - "ORDER BY time, lat, lon", - ) - xr.testing.assert_allclose(xql.to_dataset(cur, template=ds), ds.compute()) - - -def test_mysql_mixed_dimensions_register_in_a_database(mysql_con, mixed_ds): - xql.register(mysql_con, "era5", mixed_ds, table_names=NAMES) - - cur = _query( - mysql_con, - "SELECT time, level, lat, lon, temperature FROM era5.atmosphere " - "ORDER BY time, level, lat, lon", - ) - out = xql.to_dataset(cur, template=mixed_ds[["temperature"]]) - xr.testing.assert_allclose(out.temperature, mixed_ds.temperature.compute()) - default = _query(mysql_con, "SELECT DATABASE()").fetchone()[0] - assert default != "era5" + assert db.query("SELECT DATABASE()").fetchone()[0] == before From af90d89931f7c0208dd59280d5a58c76d82d1630 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 18:18:39 -0700 Subject: [PATCH 13/31] Document the tested databases, temporary tables, and ingest options Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- docs/engines.md | 46 ++++++++++++++++++++++++++++++++++----------- docs/limitations.md | 30 ++++++++++++++++++++--------- 2 files changed, 56 insertions(+), 20 deletions(-) diff --git a/docs/engines.md b/docs/engines.md index f359b9b9..1e25775b 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -239,7 +239,13 @@ Options specific to this adapter: which is how to load a long time series in slices. - `temporary=True` creates temporary tables that the database drops when the connection closes — the closest match to the other - engines' register-for-this-session behavior. + engines' register-for-this-session behavior. Where the driver cannot + create them (see the table below), registration raises rather than + risk a permanent table: Trino's driver, for one, silently ignores + the request. +- `ingest_options={...}` sets driver-specific options on each ingest + statement — e.g. Spark's staging area, + `{"spark.ingest.staging_area_uri": "s3://bucket/path"}`. - Ingest runs inside the connection's current transaction. The tables are visible to this connection at once; call `con.commit()` for other connections to see them. @@ -290,15 +296,33 @@ the Dataset (`era5.surface`), and `temporary=True` creates `Memory` tables. To choose the engine or sort key yourself, create the table first and register with `mode="append"`. -**MySQL.** Mixed-dimension Datasets go into a MySQL database named -after the Dataset, queried as `era5.surface`; identifiers are quoted -with backticks there (and on BigQuery), as those dialects require. - -**Durations.** A `timedelta64` coordinate (a forecast step, say) -round-trips everywhere, however the database stores it: as a duration -where one exists, as an integer count of its unit on SQLite and -ClickHouse, and as intervals (DuckDB, PostgreSQL) or text (MySQL) that -`to_dataset` converts back using the template. +**Tested databases.** Databases differ in how they quote identifiers, +whether they have schemas, which types they store, and what their +drivers support; the adapter keeps those facts in one table of +dialects and has a single code path. The test suite runs the same +contract — round-trips, every mode, temporary tables, mixed-dimension +naming, missing values, timedelta coordinates, the chunked round-trip — +against each database it can reach: + +| Database | `name.group` as | Temporary tables | Notes | +|---|---|---|---| +| SQLite | flat `name_group` (no schemas) | yes | times stored as text, timedeltas as integers | +| DuckDB | schema | yes | | +| PostgreSQL | schema | yes | a failed statement aborts the transaction | +| MySQL, MariaDB | database | yes | backtick identifiers; the driver ignores the target schema, so the adapter switches the default database for the ingest | +| ClickHouse, chDB | database | yes (`Memory`) | tables created by the adapter (above) | +| DataFusion | schema | no | | +| Trino | schema | no | | + +Spark, BigQuery, Databricks, and Snowflake follow their drivers' +published feature tables (no temporary tables; backtick identifiers in +Spark, BigQuery, and Databricks; no target schema in Spark, whose +groups are flat) but are not exercised by the test suite; other +databases get standard SQL. Where a database widens a type — SQLite +stores `float32` as `float64` and `bool` as an integer, MySQL `bool` as +`int8`, interval or text durations — `to_dataset` narrows a plain +`SELECT` back to the template's type; derived values such as an `AVG` +keep the result's type. The cursor is a one-shot Arrow stream: `xql.to_dataset(cur, ...)` round-trips eagerly, and `chunks=` needs `spill=True`. @@ -318,7 +342,7 @@ What each integration provides. Known issues and constraints live on | `geometry` column ([geospatial](geospatial.md#geoarrow-point-geometry-columns)) | annotated WKB passes through | native `GEOMETRY` (`"wkb"` encoding) | plain binary/struct | driver-dependent | | Mixed-dimension datasets | one schema, `name.group` tables | `name.group` views over `name_group` tables | `xql.arrow_datasets(ds, name)`, one per group | `name.group` tables in a schema; `name_group` without schemas | | Naming those tables (`table_names=`) | yes | yes | yes | yes | -| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.0` (tested on 1.12 with SQLite, DuckDB, PostgreSQL 18, MySQL 8.4, ClickHouse 26.8, and chDB 26.7 drivers) | +| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.0` (tested on 1.12; see [tested databases](#adbc-adapter-any-database-with-a-driver)) | [^spill-only]: Why DuckDB relations do not re-execute — and two other engine-specific issues worth knowing — is explained on diff --git a/docs/limitations.md b/docs/limitations.md index a043a118..8067b013 100644 --- a/docs/limitations.md +++ b/docs/limitations.md @@ -93,15 +93,27 @@ Pick your engine: - *What to do:* select the region and variables you will query before registering; use `mode="append"` to load in slices. - **SQLite has no timestamp type.** - - - *Symptom:* a datetime dimension comes back from SQLite as text. - The eager round-trip recovers it from the template, but the - chunked round-trip (`chunks=..., spill=True`) cannot build its - window predicates against a text column and fails. - - *What to do:* use the eager round-trip with SQLite, or a database - with native timestamps (PostgreSQL, DuckDB, Snowflake, ...) for - the chunked one. SQLite also widens `float32` to `float64`. + **Not every driver can create temporary tables.** + + - *Symptom:* `temporary=True` raises `ValueError` on DataFusion, + Trino, Spark, BigQuery, Databricks, and Snowflake. + - *Why:* their drivers do not support temporary ingest, and Trino's + silently creates a permanent table instead, so the adapter refuses + up front. + - *What to do:* register without `temporary=True` and drop the + tables when done. + + **Spark needs an object-store staging area.** + + - *Symptom:* ingest into Spark fails with `must set + spark.ingest.staging_area_uri`. + - *What to do:* pass + `ingest_options={"spark.ingest.staging_area_uri": "s3://..."}`; + local paths are not accepted. + + **Cloud warehouses are untested.** Snowflake, BigQuery, Databricks, + and Redshift need accounts the test suite does not have; their + handling follows the drivers' documentation. ## Constraints in any engine From 9242344956366df3f056d93e0dc9447e08e62099 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 18:23:19 -0700 Subject: [PATCH 14/31] Return a str from the test's driver-path helper (mypy) Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index ac52e8d2..c21ec8d8 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -33,7 +33,7 @@ def _module_driver(module: str) -> str | None: """The driver library an ``adbc_driver_*`` Python package ships.""" try: - return importlib.import_module(module)._driver_path() + return str(importlib.import_module(module)._driver_path()) except ImportError: return None From 2067d87201f7238bdd4656c19b16789140e955a9 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 18:24:47 -0700 Subject: [PATCH 15/31] Run the ADBC contract against database servers in CI (first pass) PostgreSQL, MySQL, MariaDB, ClickHouse, Trino, and SQL Server as service containers, with their ADBC drivers installed by dbc. SQL Server joins the contract suite here; it could not be run locally. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- .github/workflows/adbc-databases.yml | 115 +++++++++++++++++++++++++++ tests/test_adbc_backend.py | 11 ++- 2 files changed, 124 insertions(+), 2 deletions(-) create mode 100644 .github/workflows/adbc-databases.yml diff --git a/.github/workflows/adbc-databases.yml b/.github/workflows/adbc-databases.yml new file mode 100644 index 00000000..43196f8f --- /dev/null +++ b/.github/workflows/adbc-databases.yml @@ -0,0 +1,115 @@ +# Runs the ADBC adapter's contract tests against real database servers in +# service containers. The main CI job covers the in-process backends +# (SQLite, DuckDB, chDB, DataFusion); this one covers the servers. +name: adbc databases + +on: + push: + branches: [ main ] + pull_request: + paths: + - "xarray_sql/backends/adbc.py" + - "xarray_sql/backends/_adbc_dialects.py" + - "xarray_sql/ds.py" + - "xarray_sql/lazyscan.py" + - "xarray_sql/roundtrip.py" + - "tests/test_adbc_backend.py" + - ".github/workflows/adbc-databases.yml" + workflow_dispatch: + +jobs: + databases: + name: "ADBC contract on database servers" + runs-on: ubuntu-latest + services: + postgresql: + image: postgres:18 + env: + POSTGRES_PASSWORD: xql + ports: ["5432:5432"] + options: >- + --health-cmd "pg_isready -U postgres" + --health-interval 5s --health-timeout 5s --health-retries 30 + mysql: + image: mysql:8.4 + env: + MYSQL_ROOT_PASSWORD: xql + MYSQL_DATABASE: xql + ports: ["3306:3306"] + options: >- + --health-cmd "mysqladmin ping -h 127.0.0.1 -uroot -pxql" + --health-interval 5s --health-timeout 5s --health-retries 40 + mariadb: + image: mariadb:11.4 + env: + MARIADB_ROOT_PASSWORD: xql + MARIADB_DATABASE: xql + ports: ["3307:3306"] + options: >- + --health-cmd "healthcheck.sh --connect --innodb_initialized" + --health-interval 5s --health-timeout 5s --health-retries 40 + clickhouse: + image: clickhouse/clickhouse-server:26.8 + env: + CLICKHOUSE_PASSWORD: xql + ports: ["8123:8123"] + trino: + image: trinodb/trino:483 + ports: ["8080:8080"] + mssql: + image: mcr.microsoft.com/mssql/server:2022-latest + env: + ACCEPT_EULA: "Y" + MSSQL_SA_PASSWORD: XqlPassw0rd + ports: ["1433:1433"] + env: + VIRTUAL_ENV: ${{ github.workspace }}/.venv + XARRAY_SQL_TEST_POSTGRESQL_URI: postgresql://postgres:xql@localhost:5432/postgres + XARRAY_SQL_TEST_MYSQL_URI: mysql://root:xql@127.0.0.1:3306/xql + XARRAY_SQL_TEST_MARIADB_URI: mysql://root:xql@127.0.0.1:3307/xql + XARRAY_SQL_TEST_CLICKHOUSE_URI: http://localhost:8123/?user=default&password=xql + XARRAY_SQL_TEST_TRINO_URI: http://ci@localhost:8080?catalog=memory&schema=default + XARRAY_SQL_TEST_MSSQL_URI: sqlserver://sa:XqlPassw0rd@localhost:1433?database=master + steps: + - uses: actions/checkout@v4 + + - uses: dtolnay/rust-toolchain@stable + + - name: Setup sccache + uses: mozilla-actions/sccache-action@v0.0.9 + + - name: Configure sccache + run: | + echo "SCCACHE_GHA_ENABLED=true" >> $GITHUB_ENV + echo "RUSTC_WRAPPER=sccache" >> $GITHUB_ENV + + - uses: astral-sh/setup-uv@v5 + with: + python-version: "3.12" + enable-cache: true + + - name: Install xarray_sql + run: uv sync --dev --no-install-package xarray-sql + - name: build rust + run: uv run --no-project maturin develop --uv + + - name: Install ADBC drivers + run: | + uv pip install dbc adbc-driver-postgresql + for driver in chdb datafusion mysql clickhouse trino mssql; do + uv run --no-project dbc install "$driver" + done + + - name: Wait for ClickHouse and Trino + run: | + for i in $(seq 1 60); do + curl -sf http://localhost:8123/ping >/dev/null \ + && curl -sf http://localhost:8080/v1/info | grep -q '"starting":false' \ + && exit 0 + sleep 3 + done + echo "ClickHouse or Trino did not start" >&2 + exit 1 + + - name: Run the ADBC contract tests + run: uv run --no-project pytest -v -rs tests/test_adbc_backend.py diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index c21ec8d8..9b65b7a4 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -9,8 +9,8 @@ - ``chdb`` and ``datafusion`` run in-process once their drivers are installed (``dbc install chdb datafusion``). -- ``clickhouse``, ``postgresql``, ``mysql``, ``mariadb``, and ``trino`` - need a server: set ``XARRAY_SQL_TEST__URI`` to its URI (and +- ``clickhouse``, ``postgresql``, ``mysql``, ``mariadb``, ``trino``, and + ``mssql`` need a server: set ``XARRAY_SQL_TEST__URI`` to its URI (and ``XARRAY_SQL_TEST_CLICKHOUSE_DRIVER`` for a ClickHouse driver that ``dbc`` did not install). """ @@ -130,6 +130,13 @@ def connect(self): temporary=False, needs_uri=True, ), + Backend( + "mssql", + "mssql", + uri=_env("mssql"), + drop_schema="DROP SCHEMA IF EXISTS {}", + needs_uri=True, + ), ] NAMES = { From 763f3464f0de9b9dbfd350b48bd6b0547e56b5de Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 18:59:34 -0700 Subject: [PATCH 16/31] Support SQL Server in the ADBC adapter Found by running the contract against SQL Server 2022 in CI: - T-SQL has no CREATE SCHEMA IF NOT EXISTS, so mixed-dimension Datasets fell back to flat tables. Dialects can now supply their own schema DDL; SQL Server's checks SCHEMA_ID and runs CREATE SCHEMA through EXEC, as it must be alone in its batch. - SQL Server has no duration type, so timedeltas are stored as integer counts, as on SQLite and ClickHouse. - Temporary tables are queried as #name, which the tests and docs now reflect. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- docs/engines.md | 1 + tests/test_adbc_backend.py | 6 ++++- xarray_sql/backends/_adbc_dialects.py | 37 +++++++++++++++++++++++++++ xarray_sql/backends/adbc.py | 3 +-- 4 files changed, 44 insertions(+), 3 deletions(-) diff --git a/docs/engines.md b/docs/engines.md index 1e25775b..00250089 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -313,6 +313,7 @@ against each database it can reach: | ClickHouse, chDB | database | yes (`Memory`) | tables created by the adapter (above) | | DataFusion | schema | no | | | Trino | schema | no | | +| SQL Server | schema | yes (queried as `#name`) | timedeltas as integers | Spark, BigQuery, Databricks, and Snowflake follow their drivers' published feature tables (no temporary tables; backtick identifiers in diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 9b65b7a4..5b9afcc7 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -66,6 +66,8 @@ class Backend: """Whether mixed-dimension Datasets register as ``name.group``.""" temporary: bool = True """Whether ``temporary=True`` is supported.""" + temporary_prefix: str = "" + """How a temporary table's name is prefixed in queries.""" drop_schema: str = "DROP SCHEMA IF EXISTS {} CASCADE" quote: str = '"' needs_uri: bool = False @@ -135,6 +137,7 @@ def connect(self): "mssql", uri=_env("mssql"), drop_schema="DROP SCHEMA IF EXISTS {}", + temporary_prefix="#", needs_uri=True, ), ] @@ -325,7 +328,8 @@ def test_temporary_tables(db, ds): xql.register(db.con, table, ds, temporary=True) - count = db.query(f"SELECT COUNT(*) FROM {table}").fetchone()[0] + temporary = f"{db.backend.temporary_prefix}{table}" + count = db.query(f"SELECT COUNT(*) FROM {temporary}").fetchone()[0] assert count == 8 * 5 * 6 diff --git a/xarray_sql/backends/_adbc_dialects.py b/xarray_sql/backends/_adbc_dialects.py index 8e287291..e5a03de7 100644 --- a/xarray_sql/backends/_adbc_dialects.py +++ b/xarray_sql/backends/_adbc_dialects.py @@ -62,11 +62,40 @@ class Dialect: table_ddl: TableDDL | None = None """For drivers that can only append: creates each table beforehand.""" + schema_ddl: Callable[[str], str] | None = None + """For SQL without ``CREATE ... IF NOT EXISTS``: the statement that + creates the schema of a given name unless it exists.""" + def quote_identifier(self, identifier: str) -> str: """*identifier* as a quoted identifier in this database's SQL.""" escaped = identifier.replace(self.quote, self.quote * 2) return f"{self.quote}{escaped}{self.quote}" + def create_schema_sql(self, name: str) -> str: + """The statement that creates schema *name* unless it exists.""" + if self.schema_ddl is not None: + return self.schema_ddl(name) + target = self.quote_identifier(name) + return f"CREATE {self.schema_kind} IF NOT EXISTS {target}" + + +def _tsql_literal(text: str) -> str: + """*text* as a T-SQL Unicode string literal.""" + return "N'" + text.replace("'", "''") + "'" + + +def _sql_server_schema_ddl(name: str) -> str: + """Creates schema *name* unless it exists, in T-SQL. + + T-SQL has no ``CREATE SCHEMA IF NOT EXISTS``, and ``CREATE SCHEMA`` + must be alone in its batch, hence the dynamic ``EXEC``. + """ + create = f"CREATE SCHEMA {SQL_SERVER.quote_identifier(name)}" + return ( + f"IF SCHEMA_ID({_tsql_literal(name)}) IS NULL " + f"EXEC({_tsql_literal(create)})" + ) + _CLICKHOUSE_TYPES = { pa.bool_(): "Bool", @@ -167,6 +196,12 @@ def _clickhouse_ddl( table_ddl=_clickhouse_ddl, ) +SQL_SERVER = Dialect( + "SQL Server", + durations=False, + schema_ddl=_sql_server_schema_ddl, +) + DIALECTS: dict[str, Dialect] = { # Exercised by the test suite, against a live database or driver. "sqlite": Dialect("SQLite", schemas=False, durations=False), @@ -183,6 +218,8 @@ def _clickhouse_ddl( "datafusion": Dialect("DataFusion", temporary_tables=False), # Trino's driver ignores temporary=True and creates a permanent table. "trino": Dialect("Trino", temporary_tables=False), + # Temporary tables are queried as #name. + "sql server": SQL_SERVER, # From the drivers' published feature tables and the databases' SQL # references; not exercised by the test suite. "spark": Dialect("Spark", quote="`", schemas=False, temporary_tables=False), diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index 439105a3..1c75aa90 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -191,10 +191,9 @@ def _create_schema(con: dbapi.Connection, name: str, dialect: Dialect) -> bool: return _flat(name, f"{dialect.name} has no schemas to hold them") if _schema_exists(con, name): return True - target = dialect.quote_identifier(name) try: with con.cursor() as cur: - cur.execute(f"CREATE {dialect.schema_kind} IF NOT EXISTS {target}") + cur.execute(dialect.create_schema_sql(name)) except Exception as exc: if not _connection_usable(con): raise RuntimeError( From 869add69139453929aeca124783a30250a4bb973 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 19:04:51 -0700 Subject: [PATCH 17/31] Automate the per-database corner cases from review in the contract Every backend now also checks: a time filter with a plain literal, SQL keywords and spaces in variable names, quotes and non-Latin text in string coordinates, int64 and unsigned extremes, uint64 beyond the int64 range (exact or refused loudly, never wrapped), nanosecond times (kept, or a warning where the database stores microseconds), a 200k-row ingest over 10 chunks, an empty result, a 60-character table name, and float32 NaN. The ClickHouse-only time-literal test becomes the cross-database time-filter test. Some of these fail on some databases until the fixes that follow. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 174 +++++++++++++++++++++++++++++++++---- 1 file changed, 159 insertions(+), 15 deletions(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 5b9afcc7..1d74896d 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -70,6 +70,10 @@ class Backend: """How a temporary table's name is prefixed in queries.""" drop_schema: str = "DROP SCHEMA IF EXISTS {} CASCADE" quote: str = '"' + time_literal: str = "'{}'" + """How a timestamp literal is written in a comparison.""" + microseconds: bool = False + """Whether the database stores times only to the microsecond.""" needs_uri: bool = False def connect(self): @@ -107,6 +111,7 @@ def connect(self): "postgresql", _module_driver("adbc_driver_postgresql") or "postgresql", uri=_env("postgresql"), + microseconds=True, needs_uri=True, ), Backend( @@ -115,6 +120,7 @@ def connect(self): uri=_env("mysql"), drop_schema="DROP DATABASE IF EXISTS {}", quote="`", + microseconds=True, needs_uri=True, ), Backend( @@ -123,6 +129,7 @@ def connect(self): uri=_env("mariadb"), drop_schema="DROP DATABASE IF EXISTS {}", quote="`", + microseconds=True, needs_uri=True, ), Backend( @@ -130,6 +137,7 @@ def connect(self): "trino", uri=_env("trino"), temporary=False, + time_literal="TIMESTAMP '{}'", needs_uri=True, ), Backend( @@ -138,6 +146,7 @@ def connect(self): uri=_env("mssql"), drop_schema="DROP SCHEMA IF EXISTS {}", temporary_prefix="#", + microseconds=True, needs_uri=True, ), ] @@ -162,6 +171,10 @@ def name(self, base: str) -> str: self._created.append(name) return name + def quoted(self, identifier: str) -> str: + quote = self.backend.quote + return f"{quote}{identifier.replace(quote, quote * 2)}{quote}" + def query(self, sql: str): cur = self.con.cursor() cur.execute(sql) @@ -219,6 +232,7 @@ def ds() -> xr.Dataset: attrs=dict(description="Synthetic weather."), ).chunk({"time": 4}) weather["temperature"][0, 0, 0] = np.nan + weather["sst"][0, 0, 1] = np.nan return weather @@ -386,6 +400,151 @@ def test_chunked_round_trip_spills_the_cursor(db, ds): ) +def test_time_filters_select_the_right_rows(db, ds): + # A literal means UTC everywhere, and SQLite's text times compare + # with it correctly. + table = db.name("weather") + xql.register(db.con, table, ds) + literal = db.backend.time_literal.format("2021-01-01 04:00:00") + + cur = db.query( + f"SELECT time, lat, lon, temperature FROM {table} " + f"WHERE time >= {literal} ORDER BY time, lat, lon" + ) + out = xql.to_dataset(cur, template=ds) + + expected = ds.temperature.isel(time=slice(4, None)) + xr.testing.assert_allclose(out.temperature, expected.compute()) + + +def test_awkward_variable_names_round_trip(db): + ds = xr.Dataset( + { + "select": ("x", np.arange(3.0)), + "wind speed": ("x", np.arange(3.0) * 2), + "Order": ("x", np.arange(3.0) * 3), + }, + coords={"x": [10, 20, 30]}, + ).chunk({"x": 3}) + table = db.name("awkward") + xql.register(db.con, table, ds) + + columns = ", ".join( + db.quoted(n) for n in ["x", "select", "wind speed", "Order"] + ) + cur = db.query(f"SELECT {columns} FROM {table} ORDER BY {db.quoted('x')}") + out = xql.to_dataset(cur, template=ds) + + xr.testing.assert_identical(out, ds.compute()) + + +def test_text_coordinates_round_trip(db): + ds = xr.Dataset( + {"count": ("station", np.arange(4))}, + coords={"station": ["O'Hare", 'say "hi"', "東京", "Zürich"]}, + ).chunk({"station": 4}) + table = db.name("stations") + xql.register(db.con, table, ds) + + cur = db.query(f"SELECT station, count FROM {table}") + out = xql.to_dataset(cur, template=ds) + + xr.testing.assert_identical( + out.sortby("station"), ds.compute().sortby("station") + ) + + +def test_integer_extremes_round_trip(db): + info64 = np.iinfo(np.int64) + ds = xr.Dataset( + { + "i64": ("x", np.array([info64.min, 0, info64.max])), + "u8": ("x", np.array([0, 1, 255], dtype=np.uint8)), + "u32": ("x", np.array([0, 1, 2**32 - 1], dtype=np.uint32)), + "u64": ("x", np.array([0, 1, 2**62], dtype=np.uint64)), + }, + coords={"x": [0, 1, 2]}, + ).chunk({"x": 3}) + table = db.name("extremes") + xql.register(db.con, table, ds) + + cur = db.query(f"SELECT x, i64, u8, u32, u64 FROM {table} ORDER BY x") + out = xql.to_dataset(cur, template=ds) + + xr.testing.assert_identical(out, ds.compute()) + + +def test_uint64_beyond_int64_is_never_silently_wrong(db): + ds = xr.Dataset( + {"u64": ("x", np.array([2**63, 2**64 - 1], dtype=np.uint64))}, + coords={"x": [0, 1]}, + ).chunk({"x": 2}) + table = db.name("huge") + try: + xql.register(db.con, table, ds) + except (ValueError, dbapi.Error): + return # refused loudly: acceptable + + cur = db.query(f"SELECT x, u64 FROM {table} ORDER BY x") + try: + out = xql.to_dataset(cur, template=ds) + except (ValueError, TypeError, dbapi.Error): + return # refused loudly on the way back: acceptable + xr.testing.assert_identical(out, ds.compute()) + + +def test_nanosecond_times_are_kept_or_truncation_is_reported(db): + times = pd.to_datetime(["2021-01-01", "2021-01-01"]) + pd.to_timedelta( + [1, 2], unit="ns" + ) + ds = xr.Dataset({"v": ("time", [1.0, 2.0])}, coords={"time": times}).chunk( + {"time": 2} + ) + table = db.name("nanos") + if db.backend.microseconds: + with pytest.warns(RuntimeWarning, match="microsecond"): + xql.register(db.con, table, ds) + return + + xql.register(db.con, table, ds) + + cur = db.query(f"SELECT time, v FROM {table} ORDER BY time") + xr.testing.assert_identical(xql.to_dataset(cur, template=ds), ds.compute()) + + +def test_many_chunks_ingest(db): + ds = xr.Dataset( + {"v": (["time", "x"], np.arange(200_000.0).reshape(1000, 200))}, + coords={"time": np.arange(1000), "x": np.arange(200)}, + ).chunk({"time": 100}) + table = db.name("big") + xql.register(db.con, table, ds) + + count, total = db.query(f"SELECT COUNT(*), SUM(v) FROM {table}").fetchone() + assert (count, float(total)) == (200_000, float(ds.v.sum())) + + +def test_empty_result_round_trips(db, ds): + table = db.name("weather") + xql.register(db.con, table, ds) + + cur = db.query( + f"SELECT time, lat, lon, temperature FROM {table} WHERE lat > 1000" + ) + out = xql.to_dataset(cur, template=ds) + + assert out.temperature.size == 0 + + +def test_long_table_name(db, ds): + table = db.name("t" * 51) # 60 characters with the unique suffix + assert len(table) == 60 + xql.register(db.con, table, ds) + + count = db.query(f"SELECT COUNT(*) FROM {table}").fetchone()[0] + assert count == 8 * 5 * 6 + + # One database's specifics --------------------------------------------------- @@ -429,21 +588,6 @@ def test_postgresql_uses_an_existing_schema(db, mixed_ds): assert count == 6 * 3 * 4 -def test_clickhouse_time_literals_mean_utc(db, ds): - _only(db, "clickhouse", "chdb") - table = db.name("weather") - xql.register(db.con, table, ds) - - cur = db.query( - f"SELECT time, lat, lon, temperature FROM {table} " - "WHERE time >= '2021-01-01 04:00:00' ORDER BY time, lat, lon" - ) - out = xql.to_dataset(cur, template=ds) - - expected = ds.temperature.isel(time=slice(4, None)) - xr.testing.assert_allclose(out.temperature, expected.compute()) - - def test_mysql_keeps_the_default_database(db, mixed_ds): # The MySQL driver ignores the target schema, so the adapter switches # the default database for the ingest and must switch it back. From 5d3b314c1e6f59abe13072d87ecec3b71bfa2269 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 19:09:49 -0700 Subject: [PATCH 18/31] Never lose data silently on ingest: a type pass per database From review, each found by the automated corner cases: - PostgreSQL wrapped a uint64 above the int64 range to a negative number, and Trino, SQL Server, and BigQuery reject unsigned integers. Where a database has none, unsigned integers now widen to the next signed width, a uint64 above int64 raises ValueError, and to_dataset narrows them back exactly (never for negative values). - SQLite stored times as 2021-01-01T04:00:00, so a filter against '2021-01-01 04:00:00' matched every row. Times are now written in SQLite's own space-separated form, via numpy (Arrow's strftime needs a time-zone database even for zone-less times). - PostgreSQL, MySQL, MariaDB, and SQL Server keep times to the microsecond; registering a time coordinate with sub-microsecond values now warns instead of truncating silently. - PostgreSQL and DataFusion fold unquoted names to lowercase and Snowflake to uppercase; a name that would fold now warns that it must be quoted. These are one conversion pass on the ingest reader, keyed by new Dialect flags, replacing the duration-only conversion. The commit requirement after registering is now prominent in the docs. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- docs/engines.md | 31 ++++-- tests/test_adbc_backend.py | 17 +++- xarray_sql/backends/_adbc_dialects.py | 130 +++++++++++++++++++++----- xarray_sql/backends/adbc.py | 42 ++++++++- xarray_sql/ds.py | 7 +- 5 files changed, 189 insertions(+), 38 deletions(-) diff --git a/docs/engines.md b/docs/engines.md index 00250089..9e00d0f9 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -246,9 +246,12 @@ Options specific to this adapter: - `ingest_options={...}` sets driver-specific options on each ingest statement — e.g. Spark's staging area, `{"spark.ingest.staging_area_uri": "s3://bucket/path"}`. -- Ingest runs inside the connection's current transaction. The tables - are visible to this connection at once; call `con.commit()` for - other connections to see them. +- **Commit after registering.** Ingest runs inside the connection's + current transaction, as DB-API prescribes. The tables are visible to + this connection at once, but other connections — a BI tool, a + separate reader — see nothing until you call `con.commit()` (SQL + Server even blocks them on the lock), and `con.rollback()` discards + the tables. DuckDB's driver autocommits; MySQL commits DDL itself. Mixed-dimension Datasets are ingested into a database schema named after the Dataset, so `era5.surface` is the same SQL here as on @@ -306,20 +309,30 @@ against each database it can reach: | Database | `name.group` as | Temporary tables | Notes | |---|---|---|---| -| SQLite | flat `name_group` (no schemas) | yes | times stored as text, timedeltas as integers | +| SQLite | flat `name_group` (no schemas) | yes | times stored as text (`2021-01-01 04:00:00…`), so plain literals compare correctly; timedeltas and unsigned integers as integers | | DuckDB | schema | yes | | -| PostgreSQL | schema | yes | a failed statement aborts the transaction | -| MySQL, MariaDB | database | yes | backtick identifiers; the driver ignores the target schema, so the adapter switches the default database for the ingest | +| PostgreSQL | schema | yes | a failed statement aborts the transaction; times to the microsecond; mixed-case names need quotes | +| MySQL, MariaDB | database | yes | backtick identifiers; the driver ignores the target schema, so the adapter switches the default database for the ingest; times to the microsecond | | ClickHouse, chDB | database | yes (`Memory`) | tables created by the adapter (above) | -| DataFusion | schema | no | | +| DataFusion | schema | no | mixed-case names need quotes | | Trino | schema | no | | -| SQL Server | schema | yes (queried as `#name`) | timedeltas as integers | +| SQL Server | schema | yes (queried as `#name`) | timedeltas and unsigned integers as integers; times to the microsecond | Spark, BigQuery, Databricks, and Snowflake follow their drivers' published feature tables (no temporary tables; backtick identifiers in Spark, BigQuery, and Databricks; no target schema in Spark, whose groups are flat) but are not exercised by the test suite; other -databases get standard SQL. Where a database widens a type — SQLite +databases get standard SQL. + +Where a database lacks a type, the adapter converts on the way in +rather than let the driver lose data silently. Unsigned integers widen +to the next signed width where there are none (PostgreSQL would +otherwise wrap a `uint64` above the int64 range around to a negative +number); a `uint64` too large for int64 raises `ValueError` instead. +Registering a time coordinate with sub-microsecond values on a database +that stores microseconds warns, and a name the database folds +(`Weather` on PostgreSQL) warns that it must be quoted in queries. +Where a database widens a type — SQLite stores `float32` as `float64` and `bool` as an integer, MySQL `bool` as `int8`, interval or text durations — `to_dataset` narrows a plain `SELECT` back to the template's type; derived values such as an `AVG` diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 1d74896d..cbaf2b12 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -74,6 +74,8 @@ class Backend: """How a timestamp literal is written in a comparison.""" microseconds: bool = False """Whether the database stores times only to the microsecond.""" + folds: bool = False + """Whether unquoted names fold, so mixed-case ones need quotes.""" needs_uri: bool = False def connect(self): @@ -99,7 +101,7 @@ def connect(self): uri="chdb://", drop_schema="DROP DATABASE IF EXISTS {}", ), - Backend("datafusion", "datafusion", temporary=False), + Backend("datafusion", "datafusion", temporary=False, folds=True), Backend( "clickhouse", os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_DRIVER", "clickhouse"), @@ -112,6 +114,7 @@ def connect(self): _module_driver("adbc_driver_postgresql") or "postgresql", uri=_env("postgresql"), microseconds=True, + folds=True, needs_uri=True, ), Backend( @@ -545,6 +548,18 @@ def test_long_table_name(db, ds): assert count == 8 * 5 * 6 +def test_mixed_case_names_are_found_quoted(db, ds): + table = db.name("Weather") + if db.backend.folds: + with pytest.warns(RuntimeWarning, match="quote"): + xql.register(db.con, table, ds) + else: + xql.register(db.con, table, ds) + + count = db.query(f"SELECT COUNT(*) FROM {db.quoted(table)}").fetchone()[0] + assert count == 8 * 5 * 6 + + # One database's specifics --------------------------------------------------- diff --git a/xarray_sql/backends/_adbc_dialects.py b/xarray_sql/backends/_adbc_dialects.py index e5a03de7..d244e897 100644 --- a/xarray_sql/backends/_adbc_dialects.py +++ b/xarray_sql/backends/_adbc_dialects.py @@ -18,6 +18,7 @@ from collections.abc import Callable from typing import TYPE_CHECKING, Literal +import numpy as np import pyarrow as pa if TYPE_CHECKING: @@ -59,6 +60,23 @@ class Dialect: """Whether the database stores Arrow durations; if not, timedeltas are ingested as integer counts of their unit.""" + unsigned: bool = True + """Whether the database stores unsigned integers. If not, they are + widened to the next signed width, and a ``uint64`` above the int64 + range is refused rather than wrapped.""" + + timestamps_as_text: bool = False + """For databases without a time type: write times as text like + ``2021-01-01 04:00:00.000000000``, which compares correctly with a + literal such as ``'2021-01-01 04:00:00'`` (SQLite's own format).""" + + timestamp_unit: Literal["ns", "us"] = "ns" + """The finest time resolution the database stores.""" + + folds: Literal["lower", "upper"] | None = None + """How the database folds unquoted identifiers; names it would fold + must be quoted in queries.""" + table_ddl: TableDDL | None = None """For drivers that can only append: creates each table beforehand.""" @@ -199,33 +217,52 @@ def _clickhouse_ddl( SQL_SERVER = Dialect( "SQL Server", durations=False, + unsigned=False, + timestamp_unit="us", schema_ddl=_sql_server_schema_ddl, ) DIALECTS: dict[str, Dialect] = { # Exercised by the test suite, against a live database or driver. - "sqlite": Dialect("SQLite", schemas=False, durations=False), + "sqlite": Dialect( + "SQLite", + schemas=False, + durations=False, + unsigned=False, + timestamps_as_text=True, + ), "duckdb": Dialect("DuckDB"), - "postgresql": Dialect("PostgreSQL"), + # PostgreSQL wraps a uint64 above the int64 range without an error. + "postgresql": Dialect( + "PostgreSQL", unsigned=False, timestamp_unit="us", folds="lower" + ), # MariaDB's server reports itself as MySQL. "mysql": Dialect( "MySQL", quote="`", schema_kind="DATABASE", target_schema="default_database", + unsigned=False, + timestamp_unit="us", ), "clickhouse": CLICKHOUSE, - "datafusion": Dialect("DataFusion", temporary_tables=False), + "datafusion": Dialect("DataFusion", temporary_tables=False, folds="lower"), # Trino's driver ignores temporary=True and creates a permanent table. - "trino": Dialect("Trino", temporary_tables=False), + "trino": Dialect("Trino", temporary_tables=False, unsigned=False), # Temporary tables are queried as #name. "sql server": SQL_SERVER, # From the drivers' published feature tables and the databases' SQL # references; not exercised by the test suite. "spark": Dialect("Spark", quote="`", schemas=False, temporary_tables=False), - "bigquery": Dialect("BigQuery", quote="`", temporary_tables=False), + "bigquery": Dialect( + "BigQuery", + quote="`", + temporary_tables=False, + durations=False, + unsigned=False, + ), "databricks": Dialect("Databricks", quote="`", temporary_tables=False), - "snowflake": Dialect("Snowflake", temporary_tables=False), + "snowflake": Dialect("Snowflake", temporary_tables=False, folds="upper"), } """Known databases, keyed by a name their drivers' vendor names contain.""" @@ -256,24 +293,71 @@ def dialect_for(con: dbapi.Connection) -> Dialect: return CLICKHOUSE -def durations_as_integers( - reader: pa.RecordBatchReader, +_SIGNED_WIDTH = { + pa.uint8(): pa.int16(), + pa.uint16(): pa.int32(), + pa.uint32(): pa.int64(), + pa.uint64(): pa.int64(), +} + + +def _timestamps_as_text(array: pa.Array) -> pa.Array: + """Times as space-separated ISO text, at their own resolution.""" + values = np.datetime_as_string(array.to_numpy(zero_copy_only=False)) + text = np.char.replace(values, "T", " ", count=1) + return pa.array(text, pa.string(), mask=np.asarray(array.is_null())) + + +def _conversion(field: pa.Field, dialect: Dialect) -> pa.DataType | None: + """The type *field* must be ingested as, or ``None`` to keep it.""" + arrow_type = field.type + if pa.types.is_duration(arrow_type) and not dialect.durations: + return pa.int64() + if pa.types.is_unsigned_integer(arrow_type) and not dialect.unsigned: + return _SIGNED_WIDTH[arrow_type] + if pa.types.is_timestamp(arrow_type) and dialect.timestamps_as_text: + return pa.string() + return None + + +def _convert( + array: pa.Array, target: pa.DataType, dialect: Dialect +) -> pa.Array: + if pa.types.is_string(target): + return _timestamps_as_text(array) + try: + return array.cast(target) + except pa.ArrowInvalid as exc: + raise ValueError( + f"{dialect.name} cannot store {array.type} values above the " + f"{target} range, which it would otherwise wrap around; convert " + f"the variable (e.g. to float64) before registering it." + ) from exc + + +def ingestable( + reader: pa.RecordBatchReader, dialect: Dialect ) -> pa.RecordBatchReader: - """*reader* with duration columns as integer counts of their unit. + """*reader* with every column in a type *dialect*'s database stores. - The template's ``timedelta64`` unit matches the Arrow unit the scan - derived from it, so [xarray_sql.to_dataset][] reads the counts back - as the original durations. + Durations become integer counts of their unit (the template's + ``timedelta64`` unit matches it, so [xarray_sql.to_dataset][] reads + them back exactly), unsigned integers widen to the next signed width, + and times become text where there is no time type. """ - fields = [ - pa.field(f.name, pa.int64(), f.nullable, f.metadata) - if pa.types.is_duration(f.type) - else f - for f in reader.schema - ] - if all(f.type == g.type for f, g in zip(fields, reader.schema)): + targets = {i: _conversion(f, dialect) for i, f in enumerate(reader.schema)} + targets = {i: t for i, t in targets.items() if t is not None} + if not targets: return reader - schema = pa.schema(fields, metadata=reader.schema.metadata) - return pa.RecordBatchReader.from_batches( - schema, (batch.cast(schema) for batch in reader) - ) + schema = reader.schema + for i, target in targets.items(): + schema = schema.set(i, schema.field(i).with_type(target)) + + def batches(): + for batch in reader: + columns = list(batch.columns) + for i, target in targets.items(): + columns[i] = _convert(columns[i], target, dialect) + yield pa.RecordBatch.from_arrays(columns, schema=schema) + + return pa.RecordBatchReader.from_batches(schema, batches()) diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index 1c75aa90..b9fecccc 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -46,7 +46,7 @@ Dialect, IngestMode, dialect_for, - durations_as_integers, + ingestable, ) from .base import register_adapter from .pyarrow import XarrayPushdownDataset @@ -79,8 +79,7 @@ def _ingest( batches, so the source read and the database write overlap. """ reader = XarrayPushdownDataset(ds, chunks, **kwargs).scanner().to_reader() - if not dialect.durations: - reader = durations_as_integers(reader) + reader = ingestable(reader, dialect) if dialect.table_ddl is not None: for statement in dialect.table_ddl( table, @@ -164,6 +163,39 @@ def _connection_usable(con: dbapi.Connection) -> bool: return True +def _warn_on_lost_precision(ds: xr.Dataset, dialect: Dialect) -> None: + """Warn if *dialect* would truncate a time coordinate of *ds*.""" + if dialect.timestamp_unit == "ns": + return + for name, coord in ds.coords.items(): + if coord.dtype.kind != "M": + continue + values = coord.values.astype("datetime64[ns]") + if (values.astype("datetime64[us]") != values).any(): + warnings.warn( + f"{dialect.name} stores times to the microsecond, so the " + f"sub-microsecond part of {name!r} will be truncated.", + RuntimeWarning, + stacklevel=4, + ) + + +def _warn_on_folded_names(names: list[str], dialect: Dialect) -> None: + """Warn about names *dialect*'s database would fold if unquoted.""" + if dialect.folds is None: + return + fold = str.lower if dialect.folds == "lower" else str.upper + folded = [n for n in names if fold(n) != n] + if folded: + quoted = ", ".join(dialect.quote_identifier(n) for n in folded) + warnings.warn( + f"{dialect.name} folds unquoted names to {dialect.folds}case, " + f"so quote these in queries: {quoted}.", + RuntimeWarning, + stacklevel=4, + ) + + def _flat(name: str, reason: str) -> bool: """Warn that *name*'s groups become flat tables; always ``False``.""" warnings.warn( @@ -295,6 +327,10 @@ def register( f"tables; register without temporary=True and drop the " f"tables when done." ) + _warn_on_lost_precision(ds, dialect) + _warn_on_folded_names( + [name] if len(groups) <= 1 else [name, *names.values()], dialect + ) if len(groups) <= 1: _ingest( con, diff --git a/xarray_sql/ds.py b/xarray_sql/ds.py index 3803c0ef..1072480a 100644 --- a/xarray_sql/ds.py +++ b/xarray_sql/ds.py @@ -117,7 +117,7 @@ def _as_dtype(coord: xr.DataArray, dtype: np.dtype) -> xr.DataArray: return coord.astype(dtype) -_NARROWINGS = {("i", "b"), ("u", "b"), ("f", "f")} +_NARROWINGS = {("i", "b"), ("u", "b"), ("i", "u"), ("f", "f")} """(result kind, template kind) pairs a database may have widened.""" @@ -125,7 +125,8 @@ def _restore_dtype(var: xr.DataArray, dtype: np.dtype) -> xr.DataArray: """*var* as the template's *dtype*, when that loses nothing. Databases without a type widen it: SQLite stores ``float32`` as - ``float64`` and ``bool`` as an integer, MySQL ``bool`` as ``int8``. + ``float64`` and ``bool`` as an integer, MySQL ``bool`` as ``int8``, and + databases without unsigned integers store them as signed ones. A plain ``SELECT`` of such a column narrows back exactly; a derived value (an ``AVG`` of ``float32``, a ``SUM`` of ``bool``) does not, and keeps the result's dtype. Only in-memory values can be checked, so a @@ -136,6 +137,8 @@ def _restore_dtype(var: xr.DataArray, dtype: np.dtype) -> xr.DataArray: if (var.dtype.kind, dtype.kind) not in _NARROWINGS: return var values = var.values + if dtype.kind == "u" and values.size and values.min() < 0: + return var # the cast would wrap negative values around narrowed = values.astype(dtype) if not np.array_equal( narrowed.astype(values.dtype), values, equal_nan=True From adc96d34960abf538b55f07cdd835ad4b431a60f Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 19:14:40 -0700 Subject: [PATCH 19/31] Run the ADBC contract against GizmoSQL (DuckDB over Flight SQL) in CI Covers the Flight SQL driver as an ingest target. Test backends can now pass further database options, here GizmoSQL's credentials. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- .github/workflows/adbc-databases.yml | 19 +++++++++++++++--- tests/test_adbc_backend.py | 30 ++++++++++++++++++++++++---- 2 files changed, 42 insertions(+), 7 deletions(-) diff --git a/.github/workflows/adbc-databases.yml b/.github/workflows/adbc-databases.yml index 43196f8f..50036815 100644 --- a/.github/workflows/adbc-databases.yml +++ b/.github/workflows/adbc-databases.yml @@ -62,6 +62,15 @@ jobs: ACCEPT_EULA: "Y" MSSQL_SA_PASSWORD: XqlPassw0rd ports: ["1433:1433"] + # DuckDB served over Arrow Flight SQL. + gizmosql: + image: gizmodata/gizmosql:latest + env: + TLS_ENABLED: "0" + GIZMOSQL_USERNAME: xql + GIZMOSQL_PASSWORD: xql + ports: ["31337:31337"] + options: --init env: VIRTUAL_ENV: ${{ github.workspace }}/.venv XARRAY_SQL_TEST_POSTGRESQL_URI: postgresql://postgres:xql@localhost:5432/postgres @@ -70,6 +79,9 @@ jobs: XARRAY_SQL_TEST_CLICKHOUSE_URI: http://localhost:8123/?user=default&password=xql XARRAY_SQL_TEST_TRINO_URI: http://ci@localhost:8080?catalog=memory&schema=default XARRAY_SQL_TEST_MSSQL_URI: sqlserver://sa:XqlPassw0rd@localhost:1433?database=master + XARRAY_SQL_TEST_FLIGHTSQL_URI: grpc://localhost:31337 + XARRAY_SQL_TEST_FLIGHTSQL_USERNAME: xql + XARRAY_SQL_TEST_FLIGHTSQL_PASSWORD: xql steps: - uses: actions/checkout@v4 @@ -95,20 +107,21 @@ jobs: - name: Install ADBC drivers run: | - uv pip install dbc adbc-driver-postgresql + uv pip install dbc adbc-driver-postgresql adbc-driver-flightsql for driver in chdb datafusion mysql clickhouse trino mssql; do uv run --no-project dbc install "$driver" done - - name: Wait for ClickHouse and Trino + - name: Wait for ClickHouse, Trino, and GizmoSQL run: | for i in $(seq 1 60); do curl -sf http://localhost:8123/ping >/dev/null \ && curl -sf http://localhost:8080/v1/info | grep -q '"starting":false' \ + && (echo > /dev/tcp/localhost/31337) 2>/dev/null \ && exit 0 sleep 3 done - echo "ClickHouse or Trino did not start" >&2 + echo "ClickHouse, Trino, or GizmoSQL did not start" >&2 exit 1 - name: Run the ADBC contract tests diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index cbaf2b12..cc6cb202 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -9,10 +9,12 @@ - ``chdb`` and ``datafusion`` run in-process once their drivers are installed (``dbc install chdb datafusion``). -- ``clickhouse``, ``postgresql``, ``mysql``, ``mariadb``, ``trino``, and - ``mssql`` need a server: set ``XARRAY_SQL_TEST__URI`` to its URI (and +- ``clickhouse``, ``postgresql``, ``mysql``, ``mariadb``, ``trino``, + ``mssql``, and ``flightsql`` need a server: set + ``XARRAY_SQL_TEST__URI`` to its URI (and ``XARRAY_SQL_TEST_CLICKHOUSE_DRIVER`` for a ClickHouse driver that - ``dbc`` did not install). + ``dbc`` did not install; ``XARRAY_SQL_TEST_FLIGHTSQL_USERNAME`` and + ``_PASSWORD`` for a Flight SQL server such as GizmoSQL). """ import dataclasses @@ -62,6 +64,8 @@ class Backend: driver: str | None uri: str | None = None entrypoint: str | None = None + options: tuple[tuple[str, str], ...] = () + """Further database options, e.g. credentials.""" schemas: bool = True """Whether mixed-dimension Datasets register as ``name.group``.""" temporary: bool = True @@ -81,7 +85,10 @@ class Backend: def connect(self): if self.driver is None or (self.needs_uri and not self.uri): pytest.skip(f"{self.name} is not available; see module docstring") - kwargs = {"db_kwargs": {"uri": self.uri}} if self.uri else {} + db_kwargs = dict(self.options) + if self.uri: + db_kwargs["uri"] = self.uri + kwargs: dict = {"db_kwargs": db_kwargs} if db_kwargs else {} if self.entrypoint: kwargs["entrypoint"] = self.entrypoint try: @@ -92,6 +99,14 @@ def connect(self): pytest.skip(f"{self.name} driver is not installed ({exc})") +def _credentials(name: str) -> tuple[tuple[str, str], ...]: + return tuple( + (option, os.environ[f"XARRAY_SQL_TEST_{name.upper()}_{option.upper()}"]) + for option in ("username", "password") + if f"XARRAY_SQL_TEST_{name.upper()}_{option.upper()}" in os.environ + ) + + BACKENDS = [ Backend("sqlite", _module_driver("adbc_driver_sqlite"), schemas=False), Backend("duckdb", _duckdb_driver(), entrypoint="duckdb_adbc_init"), @@ -152,6 +167,13 @@ def connect(self): microseconds=True, needs_uri=True, ), + Backend( + "flightsql", + _module_driver("adbc_driver_flightsql") or "flightsql", + uri=_env("flightsql"), + options=_credentials("flightsql"), + needs_uri=True, + ), ] NAMES = { From e9051cb20e4bb622d86ec80c3a445a6873acbe17 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 20:42:57 -0700 Subject: [PATCH 20/31] Add realistic ARCO-ERA5 queries across every ADBC database An integration test registers a regional ARCO-ERA5 subset (Europe at 0.25 degrees, four 6-hourly steps, surface fields plus temperature on three pressure levels) on each backend, once per backend, and runs the queries a user would: an area-mean series, a grid-point series, derived wind speed, a threshold count, per-level means, a surface-to-850 hPa join, and a regional field round-tripped to xarray. Each answer is checked against xarray on the same subset. The backend table and db fixture move to tests/_adbc.py and conftest.py so both ADBC test modules share them. Marked integration, like the other ARCO-ERA5 tests; not yet run in CI. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/_adbc.py | 216 ++++++++++++++++++++++++++++ tests/conftest.py | 12 ++ tests/test_adbc_backend.py | 216 +--------------------------- tests/test_adbc_era5_integration.py | 185 ++++++++++++++++++++++++ 4 files changed, 416 insertions(+), 213 deletions(-) create mode 100644 tests/_adbc.py create mode 100644 tests/test_adbc_era5_integration.py diff --git a/tests/_adbc.py b/tests/_adbc.py new file mode 100644 index 00000000..7982ad94 --- /dev/null +++ b/tests/_adbc.py @@ -0,0 +1,216 @@ +"""The databases the ADBC tests run against, shared by their modules. + +``db`` (defined in ``conftest.py``) runs a test once per backend in +``BACKENDS``. SQLite and DuckDB always run. The others run when +available: + +- ``chdb`` and ``datafusion`` run in-process once their drivers are + installed (``dbc install chdb datafusion``). +- ``clickhouse``, ``postgresql``, ``mysql``, ``mariadb``, ``trino``, + ``mssql``, and ``flightsql`` need a server: set + ``XARRAY_SQL_TEST__URI`` to its URI (and + ``XARRAY_SQL_TEST_CLICKHOUSE_DRIVER`` for a ClickHouse driver that + ``dbc`` did not install; ``XARRAY_SQL_TEST_FLIGHTSQL_USERNAME`` and + ``_PASSWORD`` for a Flight SQL server such as GizmoSQL). +""" + +import dataclasses +import importlib.util +import os +import uuid + +import pytest + +try: + from adbc_driver_manager import dbapi +except ImportError: # the fixture skips; see Backend.connect + dbapi = None + + +def _module_driver(module: str) -> str | None: + """The driver library an ``adbc_driver_*`` Python package ships.""" + try: + return str(importlib.import_module(module)._driver_path()) + except ImportError: + return None + + +def _duckdb_driver() -> str | None: + """The shared library holding DuckDB's ADBC entrypoint.""" + for module in ("_duckdb", "duckdb.duckdb"): + try: + spec = importlib.util.find_spec(module) + except ModuleNotFoundError: + continue + if spec is not None and spec.origin: + return spec.origin + return None + + +def _env(name: str) -> str | None: + return os.environ.get(f"XARRAY_SQL_TEST_{name.upper()}_URI") + + +@dataclasses.dataclass(frozen=True) +class Backend: + """A database to run the contract tests against.""" + + name: str + driver: str | None + uri: str | None = None + entrypoint: str | None = None + options: tuple[tuple[str, str], ...] = () + """Further database options, e.g. credentials.""" + schemas: bool = True + """Whether mixed-dimension Datasets register as ``name.group``.""" + temporary: bool = True + """Whether ``temporary=True`` is supported.""" + temporary_prefix: str = "" + """How a temporary table's name is prefixed in queries.""" + drop_schema: str = "DROP SCHEMA IF EXISTS {} CASCADE" + quote: str = '"' + time_literal: str = "'{}'" + """How a timestamp literal is written in a comparison.""" + microseconds: bool = False + """Whether the database stores times only to the microsecond.""" + folds: bool = False + """Whether unquoted names fold, so mixed-case ones need quotes.""" + needs_uri: bool = False + + def connect(self): + if dbapi is None: + pytest.skip("adbc-driver-manager is not installed") + if self.driver is None or (self.needs_uri and not self.uri): + pytest.skip(f"{self.name} is not available; see module docstring") + db_kwargs = dict(self.options) + if self.uri: + db_kwargs["uri"] = self.uri + kwargs: dict = {"db_kwargs": db_kwargs} if db_kwargs else {} + if self.entrypoint: + kwargs["entrypoint"] = self.entrypoint + try: + return dbapi.connect(driver=self.driver, **kwargs) + except dbapi.Error as exc: + if self.needs_uri: + raise + pytest.skip(f"{self.name} driver is not installed ({exc})") + + +def _credentials(name: str) -> tuple[tuple[str, str], ...]: + return tuple( + (option, os.environ[f"XARRAY_SQL_TEST_{name.upper()}_{option.upper()}"]) + for option in ("username", "password") + if f"XARRAY_SQL_TEST_{name.upper()}_{option.upper()}" in os.environ + ) + + +BACKENDS = [ + Backend("sqlite", _module_driver("adbc_driver_sqlite"), schemas=False), + Backend("duckdb", _duckdb_driver(), entrypoint="duckdb_adbc_init"), + Backend( + "chdb", + "chdb", + uri="chdb://", + drop_schema="DROP DATABASE IF EXISTS {}", + ), + Backend("datafusion", "datafusion", temporary=False, folds=True), + Backend( + "clickhouse", + os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_DRIVER", "clickhouse"), + uri=_env("clickhouse"), + drop_schema="DROP DATABASE IF EXISTS {}", + needs_uri=True, + ), + Backend( + "postgresql", + _module_driver("adbc_driver_postgresql") or "postgresql", + uri=_env("postgresql"), + microseconds=True, + folds=True, + needs_uri=True, + ), + Backend( + "mysql", + "mysql", + uri=_env("mysql"), + drop_schema="DROP DATABASE IF EXISTS {}", + quote="`", + microseconds=True, + needs_uri=True, + ), + Backend( + "mariadb", + "mysql", + uri=_env("mariadb"), + drop_schema="DROP DATABASE IF EXISTS {}", + quote="`", + microseconds=True, + needs_uri=True, + ), + Backend( + "trino", + "trino", + uri=_env("trino"), + temporary=False, + time_literal="TIMESTAMP '{}'", + needs_uri=True, + ), + Backend( + "mssql", + "mssql", + uri=_env("mssql"), + drop_schema="DROP SCHEMA IF EXISTS {}", + temporary_prefix="#", + microseconds=True, + needs_uri=True, + ), + Backend( + "flightsql", + _module_driver("adbc_driver_flightsql") or "flightsql", + uri=_env("flightsql"), + options=_credentials("flightsql"), + needs_uri=True, + ), +] + + +class Database: + """A connection plus the unique names a test creates, dropped after.""" + + def __init__(self, backend: Backend, con) -> None: + self.backend = backend + self.con = con + self._created: list[str] = [] + + def name(self, base: str) -> str: + """A fresh table (or schema) name, dropped when the test ends.""" + name = f"{base}_{uuid.uuid4().hex[:8]}" + self._created.append(name) + return name + + def quoted(self, identifier: str) -> str: + quote = self.backend.quote + return f"{quote}{identifier.replace(quote, quote * 2)}{quote}" + + def query(self, sql: str): + cur = self.con.cursor() + cur.execute(sql) + return cur + + def cleanup(self) -> None: + postgresql = self.backend.name == "postgresql" + if postgresql: + self.con.rollback() + for name in self._created: + quoted = f"{self.backend.quote}{name}{self.backend.quote}" + for statement in ( + f"DROP TABLE IF EXISTS {quoted}", + self.backend.drop_schema.format(quoted), + ): + try: + self.query(statement).close() + except dbapi.Error: + if postgresql: + self.con.rollback() + if postgresql: + self.con.commit() diff --git a/tests/conftest.py b/tests/conftest.py index add0b59e..323bd505 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -4,6 +4,8 @@ import pandas as pd import xarray as xr +from ._adbc import BACKENDS, Database + def rand_wx(start: str, end: str) -> xr.Dataset: np.random.seed(42) @@ -148,3 +150,13 @@ def air_and_stations(): } ).chunk({"station": 3}) return air, stations + + +@pytest.fixture(params=BACKENDS, ids=[b.name for b in BACKENDS]) +def db(request): + """A connection to each available ADBC backend (see ``tests/_adbc.py``).""" + backend = request.param + database = Database(backend, backend.connect()) + yield database + database.cleanup() + database.con.close() diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index cc6cb202..459bf884 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -2,26 +2,11 @@ ``xql.register`` ingests a Dataset into any database with an ADBC driver, and ``xql.to_dataset`` rebuilds a labeled Dataset from the driver's Arrow -cursor. The contract tests below run once per backend in ``BACKENDS``; -tests of one database's specifics follow them. - -SQLite and DuckDB always run. The others run when available: - -- ``chdb`` and ``datafusion`` run in-process once their drivers are - installed (``dbc install chdb datafusion``). -- ``clickhouse``, ``postgresql``, ``mysql``, ``mariadb``, ``trino``, - ``mssql``, and ``flightsql`` need a server: set - ``XARRAY_SQL_TEST__URI`` to its URI (and - ``XARRAY_SQL_TEST_CLICKHOUSE_DRIVER`` for a ClickHouse driver that - ``dbc`` did not install; ``XARRAY_SQL_TEST_FLIGHTSQL_USERNAME`` and - ``_PASSWORD`` for a Flight SQL server such as GizmoSQL). +cursor. The contract tests below run once per backend (the ``db`` +fixture; which backends run, and how to enable more, is in +``tests/_adbc.py``); tests of one database's specifics follow them. """ -import dataclasses -import importlib.util -import os -import uuid - import numpy as np import pandas as pd import pytest @@ -32,207 +17,12 @@ dbapi = pytest.importorskip("adbc_driver_manager.dbapi") -def _module_driver(module: str) -> str | None: - """The driver library an ``adbc_driver_*`` Python package ships.""" - try: - return str(importlib.import_module(module)._driver_path()) - except ImportError: - return None - - -def _duckdb_driver() -> str | None: - """The shared library holding DuckDB's ADBC entrypoint.""" - for module in ("_duckdb", "duckdb.duckdb"): - try: - spec = importlib.util.find_spec(module) - except ModuleNotFoundError: - continue - if spec is not None and spec.origin: - return spec.origin - return None - - -def _env(name: str) -> str | None: - return os.environ.get(f"XARRAY_SQL_TEST_{name.upper()}_URI") - - -@dataclasses.dataclass(frozen=True) -class Backend: - """A database to run the contract tests against.""" - - name: str - driver: str | None - uri: str | None = None - entrypoint: str | None = None - options: tuple[tuple[str, str], ...] = () - """Further database options, e.g. credentials.""" - schemas: bool = True - """Whether mixed-dimension Datasets register as ``name.group``.""" - temporary: bool = True - """Whether ``temporary=True`` is supported.""" - temporary_prefix: str = "" - """How a temporary table's name is prefixed in queries.""" - drop_schema: str = "DROP SCHEMA IF EXISTS {} CASCADE" - quote: str = '"' - time_literal: str = "'{}'" - """How a timestamp literal is written in a comparison.""" - microseconds: bool = False - """Whether the database stores times only to the microsecond.""" - folds: bool = False - """Whether unquoted names fold, so mixed-case ones need quotes.""" - needs_uri: bool = False - - def connect(self): - if self.driver is None or (self.needs_uri and not self.uri): - pytest.skip(f"{self.name} is not available; see module docstring") - db_kwargs = dict(self.options) - if self.uri: - db_kwargs["uri"] = self.uri - kwargs: dict = {"db_kwargs": db_kwargs} if db_kwargs else {} - if self.entrypoint: - kwargs["entrypoint"] = self.entrypoint - try: - return dbapi.connect(driver=self.driver, **kwargs) - except dbapi.Error as exc: - if self.needs_uri: - raise - pytest.skip(f"{self.name} driver is not installed ({exc})") - - -def _credentials(name: str) -> tuple[tuple[str, str], ...]: - return tuple( - (option, os.environ[f"XARRAY_SQL_TEST_{name.upper()}_{option.upper()}"]) - for option in ("username", "password") - if f"XARRAY_SQL_TEST_{name.upper()}_{option.upper()}" in os.environ - ) - - -BACKENDS = [ - Backend("sqlite", _module_driver("adbc_driver_sqlite"), schemas=False), - Backend("duckdb", _duckdb_driver(), entrypoint="duckdb_adbc_init"), - Backend( - "chdb", - "chdb", - uri="chdb://", - drop_schema="DROP DATABASE IF EXISTS {}", - ), - Backend("datafusion", "datafusion", temporary=False, folds=True), - Backend( - "clickhouse", - os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_DRIVER", "clickhouse"), - uri=_env("clickhouse"), - drop_schema="DROP DATABASE IF EXISTS {}", - needs_uri=True, - ), - Backend( - "postgresql", - _module_driver("adbc_driver_postgresql") or "postgresql", - uri=_env("postgresql"), - microseconds=True, - folds=True, - needs_uri=True, - ), - Backend( - "mysql", - "mysql", - uri=_env("mysql"), - drop_schema="DROP DATABASE IF EXISTS {}", - quote="`", - microseconds=True, - needs_uri=True, - ), - Backend( - "mariadb", - "mysql", - uri=_env("mariadb"), - drop_schema="DROP DATABASE IF EXISTS {}", - quote="`", - microseconds=True, - needs_uri=True, - ), - Backend( - "trino", - "trino", - uri=_env("trino"), - temporary=False, - time_literal="TIMESTAMP '{}'", - needs_uri=True, - ), - Backend( - "mssql", - "mssql", - uri=_env("mssql"), - drop_schema="DROP SCHEMA IF EXISTS {}", - temporary_prefix="#", - microseconds=True, - needs_uri=True, - ), - Backend( - "flightsql", - _module_driver("adbc_driver_flightsql") or "flightsql", - uri=_env("flightsql"), - options=_credentials("flightsql"), - needs_uri=True, - ), -] - NAMES = { ("time", "lat", "lon"): "surface", ("time", "level", "lat", "lon"): "atmosphere", } -class Database: - """A connection plus the unique names a test creates, dropped after.""" - - def __init__(self, backend: Backend, con) -> None: - self.backend = backend - self.con = con - self._created: list[str] = [] - - def name(self, base: str) -> str: - """A fresh table (or schema) name, dropped when the test ends.""" - name = f"{base}_{uuid.uuid4().hex[:8]}" - self._created.append(name) - return name - - def quoted(self, identifier: str) -> str: - quote = self.backend.quote - return f"{quote}{identifier.replace(quote, quote * 2)}{quote}" - - def query(self, sql: str): - cur = self.con.cursor() - cur.execute(sql) - return cur - - def cleanup(self) -> None: - postgresql = self.backend.name == "postgresql" - if postgresql: - self.con.rollback() - for name in self._created: - quoted = f"{self.backend.quote}{name}{self.backend.quote}" - for statement in ( - f"DROP TABLE IF EXISTS {quoted}", - self.backend.drop_schema.format(quoted), - ): - try: - self.query(statement).close() - except dbapi.Error: - if postgresql: - self.con.rollback() - if postgresql: - self.con.commit() - - -@pytest.fixture(params=BACKENDS, ids=[b.name for b in BACKENDS]) -def db(request): - backend = request.param - database = Database(backend, backend.connect()) - yield database - database.cleanup() - database.con.close() - - @pytest.fixture def ds() -> xr.Dataset: rng = np.random.default_rng(3) diff --git a/tests/test_adbc_era5_integration.py b/tests/test_adbc_era5_integration.py new file mode 100644 index 00000000..75f68668 --- /dev/null +++ b/tests/test_adbc_era5_integration.py @@ -0,0 +1,185 @@ +"""Integration tests: realistic ARCO-ERA5 queries on every ADBC database. + +A regional subset of ARCO-ERA5 — Europe at 0.25°, four 6-hourly steps, +surface fields plus temperature on three pressure levels — is registered +on each backend the ``db`` fixture reaches (``tests/_adbc.py``), then +queried the way a user would: area means, a grid-point series, derived +wind speed, threshold counts, per-level statistics, and a join between +the surface and pressure-level tables. Every answer is checked against +xarray computing the same thing from the same subset. + +Reads anonymously from a public bucket. Excluded from the CI unit run +(``pytest -m "not integration"``); run deliberately with +``pytest -m integration tests/test_adbc_era5_integration.py``. +""" + +import numpy as np +import pytest +import xarray as xr + +import xarray_sql as xql + +from ._adbc import BACKENDS, Database + +pytestmark = pytest.mark.integration + +ERA5 = "gs://gcp-public-data-arco-era5/ar/full_37-1h-0p25deg-chunk-1.zarr-v3" +TIMES = slice("2020-07-01T00", "2020-07-01T18") +LATITUDE = slice(60, 35) # stored north to south +LONGITUDE = slice(0, 30) +LEVELS = [500, 850, 1000] + +T2M = "2m_temperature" +U10 = "10m_u_component_of_wind" +V10 = "10m_v_component_of_wind" + +NAMES = { + ("time", "latitude", "longitude"): "surface", + ("time", "level", "latitude", "longitude"): "atmosphere", +} + + +@pytest.fixture(scope="module") +def era5() -> xr.Dataset: + """The subset, read once and held in memory for every backend.""" + ds = xr.open_zarr( + ERA5, + chunks=None, + storage_options={"token": "anon"}, + consolidated=True, + ) + subset = ds[[T2M, U10, V10, "temperature"]].sel( + time=ds.time.sel(time=TIMES)[::6], + latitude=LATITUDE, + longitude=LONGITUDE, + level=LEVELS, + ) + return subset.load().chunk({"time": 1}) + + +@pytest.fixture(scope="module", params=BACKENDS, ids=[b.name for b in BACKENDS]) +def db(request, era5): + """Each backend, with the subset registered once for all its queries.""" + backend = request.param + database = Database(backend, backend.connect()) + name = database.name("era5") + xql.register(database.con, name, era5, table_names=NAMES) + if backend.schemas: + database.tables = (f"{name}.surface", f"{name}.atmosphere") + else: + database.tables = (f"{name}_surface", f"{name}_atmosphere") + yield database + database.cleanup() + database.con.close() + + +def test_area_mean_time_series(db, era5): + surface, _ = db.tables + t2m = db.quoted(T2M) + cur = db.query( + f"SELECT time, AVG({t2m}) AS {t2m} FROM {surface} " + "WHERE latitude BETWEEN 45 AND 55 AND longitude BETWEEN 5 AND 15 " + "GROUP BY time ORDER BY time" + ) + out = xql.to_dataset(cur, template=era5) + + expected = ( + era5[T2M] + .sel(latitude=slice(55, 45), longitude=slice(5, 15)) + .mean(["latitude", "longitude"]) + ) + xr.testing.assert_allclose(out[T2M], expected.compute(), rtol=1e-6) + + +def test_grid_point_series(db, era5): + surface, _ = db.tables + t2m = db.quoted(T2M) + cur = db.query( + f"SELECT time, {t2m} FROM {surface} " + "WHERE latitude = 51.5 AND longitude = 0 ORDER BY time" + ) + out = xql.to_dataset(cur, template=era5) + + expected = era5[T2M].sel(latitude=51.5, longitude=0.0, drop=True) + xr.testing.assert_allclose(out[T2M], expected.compute(), rtol=1e-6) + + +def test_max_wind_speed(db, era5): + surface, _ = db.tables + u, v = db.quoted(U10), db.quoted(V10) + cur = db.query( + f"SELECT time, MAX(SQRT({u} * {u} + {v} * {v})) AS wind " + f"FROM {surface} GROUP BY time ORDER BY time" + ) + out = xql.to_dataset(cur, template=era5) + + speed = np.hypot(era5[U10].astype("float64"), era5[V10].astype("float64")) + expected = speed.max(["latitude", "longitude"]).rename("wind") + xr.testing.assert_allclose(out["wind"], expected.compute(), rtol=1e-5) + + +def test_hot_cell_count(db, era5): + surface, _ = db.tables + t2m = db.quoted(T2M) + cur = db.query( + f"SELECT time, SUM(CASE WHEN {t2m} > 300 THEN 1 ELSE 0 END) AS hot " + f"FROM {surface} GROUP BY time ORDER BY time" + ) + out = xql.to_dataset(cur, template=era5) + + expected = (era5[T2M] > 300).sum(["latitude", "longitude"]) + np.testing.assert_array_equal(out["hot"].values, expected.values) + + +def test_mean_temperature_per_level(db, era5): + _, atmosphere = db.tables + temperature = db.quoted("temperature") + cur = db.query( + f"SELECT level, AVG({temperature}) AS {temperature} " + f"FROM {atmosphere} GROUP BY level ORDER BY level" + ) + out = xql.to_dataset(cur, template=era5) + + expected = era5["temperature"].mean(["time", "latitude", "longitude"]) + xr.testing.assert_allclose( + out["temperature"], expected.compute(), rtol=1e-6 + ) + + +def test_surface_to_850_hpa_join(db, era5): + surface, atmosphere = db.tables + t2m = db.quoted(T2M) + temperature = db.quoted("temperature") + cur = db.query( + f"SELECT s.time, AVG(s.{t2m} - a.{temperature}) AS difference " + f"FROM {surface} s JOIN {atmosphere} a " + "ON s.time = a.time AND s.latitude = a.latitude " + "AND s.longitude = a.longitude " + "WHERE a.level = 850 GROUP BY s.time ORDER BY s.time" + ) + out = xql.to_dataset(cur, template=era5) + + difference = era5[T2M].astype("float64") - era5["temperature"].sel( + level=850, drop=True + ).astype("float64") + expected = difference.mean(["latitude", "longitude"]).rename("difference") + xr.testing.assert_allclose(out["difference"], expected.compute(), rtol=1e-6) + + +def test_regional_field_round_trips(db, era5): + surface, _ = db.tables + t2m = db.quoted(T2M) + cur = db.query( + f"SELECT time, latitude, longitude, {t2m} FROM {surface} " + "WHERE latitude BETWEEN 50 AND 52 AND longitude BETWEEN 0 AND 2 " + "ORDER BY time, latitude, longitude" + ) + out = xql.to_dataset(cur, template=era5) + + expected = era5[T2M].sel(latitude=slice(52, 50), longitude=slice(0, 2)) + xr.testing.assert_allclose( + out[T2M].sortby("latitude", ascending=False), + expected.compute(), + ) + assert out[T2M].dtype == era5[T2M].dtype + assert out[T2M].attrs == era5[T2M].attrs From dbd9dac60d7c9949126ad16f43e6b03549830270 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 20:44:19 -0700 Subject: [PATCH 21/31] Infer result dims from a mixed-dimension template to_dataset took the template's dims from its first data variable, assuming every variable shared them. With a mixed-dimension template (ARCO-ERA5: surface fields and pressure-level fields), a result grouped by level lost that dim and failed with "no template dimension survives". The dims now come from the template variables the result carries (or from every variable, for results of aliases only), in their axis order. Found by the ARCO-ERA5 integration test. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_duckdb_backend.py | 18 ++++++++++++++++++ xarray_sql/ds.py | 32 +++++++++++++++++++------------- xarray_sql/roundtrip.py | 4 ++-- 3 files changed, 39 insertions(+), 15 deletions(-) diff --git a/tests/test_duckdb_backend.py b/tests/test_duckdb_backend.py index 26f9204b..11b9d25e 100644 --- a/tests/test_duckdb_backend.py +++ b/tests/test_duckdb_backend.py @@ -133,6 +133,24 @@ def test_to_dataset_accepts_plain_arrow_table(ds): np.testing.assert_allclose(out["temperature"].values, [1.0, 2.0, 3.0]) +def test_to_dataset_infers_dims_from_a_mixed_dimension_template(): + # Grouping a pressure-level variable by level keeps `level`, although + # the template's first variable (a surface field) has no such dim. + template = xr.Dataset( + { + "t2m": (["time", "lat"], np.zeros((2, 3))), + "temperature": (["time", "level", "lat"], np.zeros((2, 2, 3))), + }, + coords={"time": [0, 1], "level": [500, 850], "lat": [1.0, 2.0, 3.0]}, + ) + result = pa.table({"level": [500, 850], "temperature": [250.0, 280.0]}) + + out = xql.to_dataset(result, template=template) + + assert out.temperature.dims == ("level",) + np.testing.assert_array_equal(out.temperature.values, [250.0, 280.0]) + + def test_to_dataset_requires_dims_or_template(): table = pa.table({"a": [1, 2], "b": [3.0, 4.0]}) with pytest.raises(ValueError, match="dims cannot be inferred"): diff --git a/xarray_sql/ds.py b/xarray_sql/ds.py index 1072480a..6a03deaa 100644 --- a/xarray_sql/ds.py +++ b/xarray_sql/ds.py @@ -56,18 +56,24 @@ # --------------------------------------------------------------------------- -def _ds_var_dims(ds: xr.Dataset) -> list[str]: - """Return a Dataset's data-variable dim order. - - The forward path validates that all data variables share the same dims - tuple, so the first var's dim order is canonical. Falls back to - ``ds.dims`` keys for empty Datasets. Always use this rather than - ``list(ds.dims)`` when round-tripping, since the latter is in - canonical name order and may not match the variable's axis order. +def _result_dims(template: xr.Dataset, columns) -> list[str]: + """The template's dims that survive into a result's *columns*. + + In the variables' axis order, not ``template.dims``'s name order. A + mixed-dimension template has one order per group of variables, so + the variables the result carries pick theirs: ``AVG(temperature) ... + GROUP BY level`` keeps ``level`` although the template's first + variable has none. A result carrying no template variable (only + aliases such as ``wind``) may keep any template dim. """ - if ds.data_vars: - return list(next(iter(ds.data_vars.values())).dims) - return list(ds.dims) + columns = set(columns) + carried = [name for name in template.data_vars if name in columns] + order: list = [] + for name in carried or list(template.data_vars): + order.extend(d for d in template[name].dims if d not in order) + if not order: + order = list(template.dims) + return [d for d in order if d in columns] _TIMEDELTA_PARTS = ( @@ -1175,13 +1181,13 @@ def _infer_dimension_columns( become the dimensions, so aggregations that drop dims (e.g. ``GROUP BY time`` over a ``(time, lat, lon)`` grid) round-trip on the surviving dim(s). Uses the data variable's dim order (via - ``_ds_var_dims``) so the original axis order is preserved. + ``_result_dims``) so the original axis order is preserved. """ result_cols = set(self._result_columns()) def surviving(template: xr.Dataset) -> list[str]: # Template dims still present in the result, in var axis order. - return [d for d in _ds_var_dims(template) if d in result_cols] + return _result_dims(template, result_cols) if preferred_template is not None: preferred = surviving(preferred_template) diff --git a/xarray_sql/roundtrip.py b/xarray_sql/roundtrip.py index 55b3402e..cc7de61e 100644 --- a/xarray_sql/roundtrip.py +++ b/xarray_sql/roundtrip.py @@ -40,8 +40,8 @@ XarrayDataFrame, _build_lazy_scan, _dataset_from_batches, - _ds_var_dims, _finish_dataset, + _result_dims, ) from .lazyscan import LazyResultHandle, PolarsHandle, resolve_lazy_handle @@ -347,7 +347,7 @@ def _resolve_dims( "dims cannot be inferred without a template; pass " "dims=[...] or template=." ) - dims = [d for d in _ds_var_dims(template) if d in field_names] + dims = _result_dims(template, field_names) if not dims: raise ValueError( "dims cannot be inferred: no template dimension survives " From d8d32079ff46d295580e2703605a5b9fb0f9955c Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 20:47:40 -0700 Subject: [PATCH 22/31] Run the ARCO-ERA5 queries on every database in the adbc databases workflow Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- .github/workflows/adbc-databases.yml | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/.github/workflows/adbc-databases.yml b/.github/workflows/adbc-databases.yml index 50036815..1da9d823 100644 --- a/.github/workflows/adbc-databases.yml +++ b/.github/workflows/adbc-databases.yml @@ -13,7 +13,9 @@ on: - "xarray_sql/ds.py" - "xarray_sql/lazyscan.py" - "xarray_sql/roundtrip.py" + - "tests/_adbc.py" - "tests/test_adbc_backend.py" + - "tests/test_adbc_era5_integration.py" - ".github/workflows/adbc-databases.yml" workflow_dispatch: @@ -126,3 +128,9 @@ jobs: - name: Run the ADBC contract tests run: uv run --no-project pytest -v -rs tests/test_adbc_backend.py + + - name: Run realistic ARCO-ERA5 queries on every database + # Reads a regional subset anonymously from the public bucket. + run: >- + uv run --no-project pytest -v -rs -m integration + tests/test_adbc_era5_integration.py From 8cb75a94ed6ac3c273758aeac4be437c55175318 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 20:59:46 -0700 Subject: [PATCH 23/31] Write SQLite's whole-second times exactly as datetime() does SQLite times were written with a fraction of a second at their full resolution (2021-01-01 23:00:00.000000000), which sorts after the literal '2021-01-01 23:00:00', so an inclusive upper bound (BETWEEN, <=) or an equality missed the row at that time. Whole seconds are now written as SQLite's datetime() writes them; a fraction is kept only where there is one. The time-filter contract test now uses an inclusive BETWEEN. Found while wiring the geospatial benchmarks to ADBC databases. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 11 +++++++---- xarray_sql/backends/_adbc_dialects.py | 16 +++++++++++++--- 2 files changed, 20 insertions(+), 7 deletions(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 459bf884..b707c535 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -217,18 +217,21 @@ def test_chunked_round_trip_spills_the_cursor(db, ds): def test_time_filters_select_the_right_rows(db, ds): # A literal means UTC everywhere, and SQLite's text times compare - # with it correctly. + # with it correctly, including at an inclusive bound. table = db.name("weather") xql.register(db.con, table, ds) - literal = db.backend.time_literal.format("2021-01-01 04:00:00") + start, end = ( + db.backend.time_literal.format(f"2021-01-01 0{hour}:00:00") + for hour in (4, 6) + ) cur = db.query( f"SELECT time, lat, lon, temperature FROM {table} " - f"WHERE time >= {literal} ORDER BY time, lat, lon" + f"WHERE time BETWEEN {start} AND {end} ORDER BY time, lat, lon" ) out = xql.to_dataset(cur, template=ds) - expected = ds.temperature.isel(time=slice(4, None)) + expected = ds.temperature.isel(time=slice(4, 7)) xr.testing.assert_allclose(out.temperature, expected.compute()) diff --git a/xarray_sql/backends/_adbc_dialects.py b/xarray_sql/backends/_adbc_dialects.py index d244e897..9ebffa06 100644 --- a/xarray_sql/backends/_adbc_dialects.py +++ b/xarray_sql/backends/_adbc_dialects.py @@ -302,9 +302,19 @@ def dialect_for(con: dbapi.Connection) -> Dialect: def _timestamps_as_text(array: pa.Array) -> pa.Array: - """Times as space-separated ISO text, at their own resolution.""" - values = np.datetime_as_string(array.to_numpy(zero_copy_only=False)) - text = np.char.replace(values, "T", " ", count=1) + """Times as space-separated ISO text. + + Whole seconds are written exactly as SQLite's ``datetime()`` writes + them (``2021-01-01 04:00:00``), so equality and inclusive bounds + against such a literal hold; a fraction of a second is kept where + there is one, which still orders correctly. + """ + values = array.to_numpy(zero_copy_only=False) + whole = np.datetime_as_string(values, unit="s") + exact = np.datetime_as_string(values) + fractional = values != values.astype("datetime64[s]") + text = np.where(fractional, exact, whole) + text = np.char.replace(text, "T", " ", count=1) return pa.array(text, pa.string(), mask=np.asarray(array.is_null())) From 229f38a7502d155a45df0cf4c01ca17389570f4c Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 21:16:08 -0700 Subject: [PATCH 24/31] Parse spilled text times as ISO 8601, not an inferred format SQLite times are written with a fraction of a second only where there is one, so a column can mix '2021-01-01 04:00:00' and '2021-01-01 04:00:00.500000000'. The spill read parsed them with pd.to_datetime and no format, which infers one from the first value and raises on the other form, failing the chunked round-trip of any sub-second time coordinate on SQLite. A contract test now round-trips whole and half-second times, eagerly and chunked, on every backend. Found in code review. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 29 +++++++++++++++++++++++++++++ xarray_sql/roundtrip.py | 5 ++++- 2 files changed, 33 insertions(+), 1 deletion(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index b707c535..1130d821 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -235,6 +235,35 @@ def test_time_filters_select_the_right_rows(db, ds): xr.testing.assert_allclose(out.temperature, expected.compute()) +@pytest.mark.parametrize( + "chunks", [None, {"time": 2}], ids=["eager", "chunked"] +) +def test_subsecond_times_round_trip(db, chunks): + # Whole and fractional seconds in one column: SQLite writes them as + # text of two lengths, which the spilled read must parse alike. + times = pd.to_datetime( + [ + "2021-01-01 04:00:00", + "2021-01-01 04:00:00.5", + "2021-01-01 05:00:00", + "2021-01-01 05:00:00.5", + ], + format="ISO8601", + ) + ds = xr.Dataset( + {"v": ("time", [1.0, 2.0, 3.0, 4.0])}, coords={"time": times} + ).chunk({"time": 4}) + table = db.name("subsecond") + xql.register(db.con, table, ds) + + cur = db.query(f"SELECT time, v FROM {table} ORDER BY time") + out = xql.to_dataset( + cur, template=ds, chunks=chunks, spill=chunks is not None + ) + + xr.testing.assert_identical(out.compute(), ds.compute()) + + def test_awkward_variable_names_round_trip(db): ds = xr.Dataset( { diff --git a/xarray_sql/roundtrip.py b/xarray_sql/roundtrip.py index cc7de61e..a592ccc2 100644 --- a/xarray_sql/roundtrip.py +++ b/xarray_sql/roundtrip.py @@ -513,7 +513,10 @@ def _as_durations(array: pa.Array) -> pa.Array: def _as_timestamps(array: pa.Array) -> pa.Array: """Times a database returned as text (SQLite has no time type).""" - return pa.array(pd.to_datetime(array.to_pandas()), pa.timestamp("ns")) + # ISO8601 rather than an inferred format: SQLite text has a fraction + # of a second only where there is one, so a column mixes both forms. + times = pd.to_datetime(array.to_pandas(), format="ISO8601") + return pa.array(times, pa.timestamp("ns")) def _stream_to_parquet( From 324442f79c760162557575ae2b221520a6270177 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sat, 26 Sep 2026 22:48:14 -0700 Subject: [PATCH 25/31] ANALYZE tables after ingest on PostgreSQL PostgreSQL gathers statistics only in the background after a commit, so a freshly registered table has none and the planner guesses. In the geospatial benchmarks on real ARCO-ERA5 data, a climatology self-join and a raster-by-region range join over registered tables timed out after 45 minutes; with ANALYZE after ingest they finish in 5 and 14 seconds. The adapter now analyzes each table it ingests on PostgreSQL (a Dialect flag), and a test checks the statistics exist. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 11 +++++++++++ xarray_sql/backends/_adbc_dialects.py | 12 +++++++++++- xarray_sql/backends/adbc.py | 6 ++++++ 3 files changed, 28 insertions(+), 1 deletion(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 1130d821..8a8d4606 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -447,6 +447,17 @@ def test_postgresql_uses_an_existing_schema(db, mixed_ds): assert count == 6 * 3 * 4 +def test_postgresql_tables_have_statistics_after_register(db, ds): + # Without them the planner guesses, and joins over a freshly + # registered table can pick plans that run for hours. + _only(db, "postgresql") + table = db.name("weather") + xql.register(db.con, table, ds) + + cur = db.query(f"SELECT reltuples FROM pg_class WHERE relname = '{table}'") + assert cur.fetchone()[0] == 8 * 5 * 6 + + def test_mysql_keeps_the_default_database(db, mixed_ds): # The MySQL driver ignores the target schema, so the adapter switches # the default database for the ingest and must switch it back. diff --git a/xarray_sql/backends/_adbc_dialects.py b/xarray_sql/backends/_adbc_dialects.py index 9ebffa06..4118ff1a 100644 --- a/xarray_sql/backends/_adbc_dialects.py +++ b/xarray_sql/backends/_adbc_dialects.py @@ -77,6 +77,12 @@ class Dialect: """How the database folds unquoted identifiers; names it would fold must be quoted in queries.""" + analyze: bool = False + """Whether to ``ANALYZE`` a table after ingest. PostgreSQL gathers + statistics only in the background, after a commit; until then its + planner guesses, and a join over a freshly registered table can pick + a nested loop that runs for hours instead of seconds.""" + table_ddl: TableDDL | None = None """For drivers that can only append: creates each table beforehand.""" @@ -234,7 +240,11 @@ def _clickhouse_ddl( "duckdb": Dialect("DuckDB"), # PostgreSQL wraps a uint64 above the int64 range without an error. "postgresql": Dialect( - "PostgreSQL", unsigned=False, timestamp_unit="us", folds="lower" + "PostgreSQL", + unsigned=False, + timestamp_unit="us", + folds="lower", + analyze=True, ), # MariaDB's server reports itself as MySQL. "mysql": Dialect( diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index b9fecccc..ffdae3cf 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -112,6 +112,12 @@ def _ingest( db_schema_name=target_schema, temporary=temporary, ) + if dialect.analyze: + target = dialect.quote_identifier(table) + if db_schema_name is not None: + target = f"{dialect.quote_identifier(db_schema_name)}.{target}" + with con.cursor() as cur: + cur.execute(f"ANALYZE {target}") @contextlib.contextmanager From 673432c0d2773a7541fb5fde7805de54fd10cef3 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sun, 27 Sep 2026 03:46:29 -0700 Subject: [PATCH 26/31] Document MariaDB's nested-loop joins, Trino ingest, and PostgreSQL ANALYZE Found running the geospatial benchmarks on real ERA5 data: MariaDB needs hash joins enabled for joins on computed keys, and Trino's driver ingests about 10k rows/s into a memory catalog capped at 128 MB by default. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- docs/engines.md | 6 +++--- docs/limitations.md | 18 ++++++++++++++++++ 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/docs/engines.md b/docs/engines.md index 9e00d0f9..28e7e4ce 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -311,11 +311,11 @@ against each database it can reach: |---|---|---|---| | SQLite | flat `name_group` (no schemas) | yes | times stored as text (`2021-01-01 04:00:00…`), so plain literals compare correctly; timedeltas and unsigned integers as integers | | DuckDB | schema | yes | | -| PostgreSQL | schema | yes | a failed statement aborts the transaction; times to the microsecond; mixed-case names need quotes | -| MySQL, MariaDB | database | yes | backtick identifiers; the driver ignores the target schema, so the adapter switches the default database for the ingest; times to the microsecond | +| PostgreSQL | schema | yes | tables are `ANALYZE`d after ingest; a failed statement aborts the transaction; times to the microsecond; mixed-case names need quotes | +| MySQL, MariaDB | database | yes | backtick identifiers; the driver ignores the target schema, so the adapter switches the default database for the ingest; times to the microsecond; MariaDB joins without hash joins by default (see limitations) | | ClickHouse, chDB | database | yes (`Memory`) | tables created by the adapter (above) | | DataFusion | schema | no | mixed-case names need quotes | -| Trino | schema | no | | +| Trino | schema | no | ingest is slow, about 10k rows/s (see limitations) | | SQL Server | schema | yes (queried as `#name`) | timedeltas and unsigned integers as integers; times to the microsecond | Spark, BigQuery, Databricks, and Snowflake follow their drivers' diff --git a/docs/limitations.md b/docs/limitations.md index 8067b013..54473e22 100644 --- a/docs/limitations.md +++ b/docs/limitations.md @@ -111,6 +111,24 @@ Pick your engine: `ingest_options={"spark.ingest.staging_area_uri": "s3://..."}`; local paths are not accepted. + **MariaDB joins by nested loop by default.** + + - *Symptom:* a join on a computed key (e.g. `ON e.time = f.time + + f.lead`) runs for hours on MariaDB; MySQL answers the same query in + seconds with a hash join. + - *What to do:* enable MariaDB's hash joins for the session, + `SET SESSION join_cache_level = 8`, before querying. + + **Trino ingests slowly.** + + - *Symptom:* registering on Trino takes minutes per million rows. + - *Why:* its ADBC driver ingests with one `INSERT` per batch, about + 10,000 rows a second, and Trino's `memory` catalog caps stored data + (128 MB by default, `memory.max-data-per-node`). + - *What to do:* register a region or period rather than a global + field, or load large data through a Trino connector that writes + files (Hive, Iceberg). + **Cloud warehouses are untested.** Snowflake, BigQuery, Databricks, and Redshift need accounts the test suite does not have; their handling follows the drivers' documentation. From 8aee30a10fca9c138e0724f1f9bf4dc361771338 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Sun, 27 Sep 2026 13:32:39 -0700 Subject: [PATCH 27/31] Run the database CI as one job per database; cancel superseded runs The single job pulled seven service images one after another (about two minutes, SQL Server and Trino alone about 2.5 GB) before running every database's tests in sequence. It is now a matrix: each job starts only its own server and runs only its own backend (XARRAY_SQL_TEST_ONLY, a new filter in the tests' backend table), so the databases run in parallel and a failure names its database. An in-process job runs the ARCO-ERA5 queries on SQLite, DuckDB, chDB, and DataFusion, whose contract tests run in the main CI. A newer push to the same pull request now cancels the running check. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- .github/workflows/adbc-databases.yml | 141 ++++++++++++++++----------- tests/_adbc.py | 6 ++ 2 files changed, 88 insertions(+), 59 deletions(-) diff --git a/.github/workflows/adbc-databases.yml b/.github/workflows/adbc-databases.yml index 1da9d823..c27c239f 100644 --- a/.github/workflows/adbc-databases.yml +++ b/.github/workflows/adbc-databases.yml @@ -1,6 +1,9 @@ -# Runs the ADBC adapter's contract tests against real database servers in -# service containers. The main CI job covers the in-process backends -# (SQLite, DuckDB, chDB, DataFusion); this one covers the servers. +# Runs the ADBC adapter's contract tests and the ARCO-ERA5 queries against +# real database servers in service containers, one job per database so +# each pulls only its own image and a failure names its database. The +# main CI job runs the contract on the in-process backends (SQLite, +# DuckDB, chDB, DataFusion); the "in-process" job here adds their +# ARCO-ERA5 queries. name: adbc databases on: @@ -19,69 +22,70 @@ on: - ".github/workflows/adbc-databases.yml" workflow_dispatch: +# A newer push to the same pull request supersedes a running check. +concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: ${{ github.event_name == 'pull_request' }} + jobs: databases: - name: "ADBC contract on database servers" + name: ${{ matrix.backend }} runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - backend: sqlite,duckdb,chdb,datafusion + - backend: postgresql + image: postgres:18 + port: 5432 + uri: postgresql://postgres:xql@localhost:5432/postgres + - backend: mysql + image: mysql:8.4 + port: 3306 + uri: mysql://root:xql@127.0.0.1:3306/xql + - backend: mariadb + image: mariadb:11.4 + port: 3306 + uri: mysql://root:xql@127.0.0.1:3306/xql + - backend: clickhouse + image: clickhouse/clickhouse-server:26.8 + port: 8123 + uri: http://localhost:8123/?user=default&password=xql + - backend: trino + image: trinodb/trino:483 + port: 8080 + uri: http://ci@localhost:8080?catalog=memory&schema=default + - backend: mssql + image: mcr.microsoft.com/mssql/server:2022-latest + port: 1433 + uri: sqlserver://sa:XqlPassw0rd@localhost:1433?database=master + - backend: flightsql + image: gizmodata/gizmosql:latest + port: 31337 + uri: grpc://localhost:31337 services: - postgresql: - image: postgres:18 + # An empty image (the in-process databases) starts no container. + # Each image reads only its own variables below. + db: + image: ${{ matrix.image }} + ports: + - ${{ matrix.port || 1 }}:${{ matrix.port || 1 }} env: POSTGRES_PASSWORD: xql - ports: ["5432:5432"] - options: >- - --health-cmd "pg_isready -U postgres" - --health-interval 5s --health-timeout 5s --health-retries 30 - mysql: - image: mysql:8.4 - env: MYSQL_ROOT_PASSWORD: xql MYSQL_DATABASE: xql - ports: ["3306:3306"] - options: >- - --health-cmd "mysqladmin ping -h 127.0.0.1 -uroot -pxql" - --health-interval 5s --health-timeout 5s --health-retries 40 - mariadb: - image: mariadb:11.4 - env: MARIADB_ROOT_PASSWORD: xql MARIADB_DATABASE: xql - ports: ["3307:3306"] - options: >- - --health-cmd "healthcheck.sh --connect --innodb_initialized" - --health-interval 5s --health-timeout 5s --health-retries 40 - clickhouse: - image: clickhouse/clickhouse-server:26.8 - env: CLICKHOUSE_PASSWORD: xql - ports: ["8123:8123"] - trino: - image: trinodb/trino:483 - ports: ["8080:8080"] - mssql: - image: mcr.microsoft.com/mssql/server:2022-latest - env: ACCEPT_EULA: "Y" MSSQL_SA_PASSWORD: XqlPassw0rd - ports: ["1433:1433"] - # DuckDB served over Arrow Flight SQL. - gizmosql: - image: gizmodata/gizmosql:latest - env: TLS_ENABLED: "0" GIZMOSQL_USERNAME: xql GIZMOSQL_PASSWORD: xql - ports: ["31337:31337"] - options: --init env: VIRTUAL_ENV: ${{ github.workspace }}/.venv - XARRAY_SQL_TEST_POSTGRESQL_URI: postgresql://postgres:xql@localhost:5432/postgres - XARRAY_SQL_TEST_MYSQL_URI: mysql://root:xql@127.0.0.1:3306/xql - XARRAY_SQL_TEST_MARIADB_URI: mysql://root:xql@127.0.0.1:3307/xql - XARRAY_SQL_TEST_CLICKHOUSE_URI: http://localhost:8123/?user=default&password=xql - XARRAY_SQL_TEST_TRINO_URI: http://ci@localhost:8080?catalog=memory&schema=default - XARRAY_SQL_TEST_MSSQL_URI: sqlserver://sa:XqlPassw0rd@localhost:1433?database=master - XARRAY_SQL_TEST_FLIGHTSQL_URI: grpc://localhost:31337 + XARRAY_SQL_TEST_ONLY: ${{ matrix.backend }} XARRAY_SQL_TEST_FLIGHTSQL_USERNAME: xql XARRAY_SQL_TEST_FLIGHTSQL_PASSWORD: xql steps: @@ -114,22 +118,41 @@ jobs: uv run --no-project dbc install "$driver" done - - name: Wait for ClickHouse, Trino, and GizmoSQL + - name: Point the tests' backend table at the server + if: matrix.uri run: | - for i in $(seq 1 60); do - curl -sf http://localhost:8123/ping >/dev/null \ - && curl -sf http://localhost:8080/v1/info | grep -q '"starting":false' \ - && (echo > /dev/tcp/localhost/31337) 2>/dev/null \ - && exit 0 - sleep 3 - done - echo "ClickHouse, Trino, or GizmoSQL did not start" >&2 - exit 1 + echo "XARRAY_SQL_TEST_${BACKEND^^}_URI=${{ matrix.uri }}" >> "$GITHUB_ENV" + env: + BACKEND: ${{ matrix.backend }} + + - name: Wait for the server + if: matrix.uri + run: | + uv run --no-project python - <<'PY' + import os, sys, time + sys.path.insert(0, os.getcwd()) + from tests._adbc import BACKENDS + backend = next(b for b in BACKENDS if b.name == "${{ matrix.backend }}") + for attempt in range(100): + try: + con = backend.connect() + cur = con.cursor() + cur.execute("SELECT 1") + cur.fetchall() + print(f"{backend.name} is up") + break + except Exception as exc: + print(f"waiting for {backend.name}: {str(exc)[:120]}") + time.sleep(3) + else: + sys.exit(f"{backend.name} did not start") + PY - name: Run the ADBC contract tests + if: matrix.uri run: uv run --no-project pytest -v -rs tests/test_adbc_backend.py - - name: Run realistic ARCO-ERA5 queries on every database + - name: Run realistic ARCO-ERA5 queries # Reads a regional subset anonymously from the public bucket. run: >- uv run --no-project pytest -v -rs -m integration diff --git a/tests/_adbc.py b/tests/_adbc.py index 7982ad94..dd866015 100644 --- a/tests/_adbc.py +++ b/tests/_adbc.py @@ -12,6 +12,9 @@ ``XARRAY_SQL_TEST_CLICKHOUSE_DRIVER`` for a ClickHouse driver that ``dbc`` did not install; ``XARRAY_SQL_TEST_FLIGHTSQL_USERNAME`` and ``_PASSWORD`` for a Flight SQL server such as GizmoSQL). + +``XARRAY_SQL_TEST_ONLY`` (comma-separated names) restricts a run to +those backends, e.g. one CI job per database. """ import dataclasses @@ -80,6 +83,9 @@ class Backend: def connect(self): if dbapi is None: pytest.skip("adbc-driver-manager is not installed") + only = os.environ.get("XARRAY_SQL_TEST_ONLY") + if only and self.name not in only.split(","): + pytest.skip(f"XARRAY_SQL_TEST_ONLY={only} excludes {self.name}") if self.driver is None or (self.needs_uri and not self.uri): pytest.skip(f"{self.name} is not available; see module docstring") db_kwargs = dict(self.options) From 8a6108c68df37c43fc00d2dcfaad1566842fa0e6 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Mon, 28 Sep 2026 08:34:33 -0700 Subject: [PATCH 28/31] Restore a widened dtype only for plain selects, decided by dims to_dataset narrowed a result column back to the template variable's dtype whenever the values survived exactly, so SUM(flag) AS flag came back as bool when every sum was 0 or 1 and as an integer otherwise: the same query, a dtype depending on the data. The query is not available at this point, but its shape is: a plain select (filtered or not) keeps every dim of the variable, while an aggregate reduces some away. Narrowing now applies only when the column keeps all of the variable's dims. From review. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_duckdb_backend.py | 17 +++++++++++++++++ xarray_sql/ds.py | 24 ++++++++++++++++-------- 2 files changed, 33 insertions(+), 8 deletions(-) diff --git a/tests/test_duckdb_backend.py b/tests/test_duckdb_backend.py index 11b9d25e..e592cfc5 100644 --- a/tests/test_duckdb_backend.py +++ b/tests/test_duckdb_backend.py @@ -151,6 +151,23 @@ def test_to_dataset_infers_dims_from_a_mixed_dimension_template(): np.testing.assert_array_equal(out.temperature.values, [250.0, 280.0]) +def test_widened_dtype_is_restored_for_plain_selects_only(): + # SQLite and MySQL return bool as an integer. A plain select of the + # variable narrows back; an aggregate aliased to the variable's name + # keeps its integer dtype whatever its values happen to be. + template = xr.Dataset( + {"flag": (["time", "x"], np.array([[True, False], [False, False]]))}, + coords={"time": [0, 1], "x": [10, 20]}, + ) + plain = pa.table( + {"time": [0, 0, 1, 1], "x": [10, 20, 10, 20], "flag": [1, 0, 0, 0]} + ) + summed = pa.table({"x": [10, 20], "flag": [1, 0]}) # SUM(flag) by x + + assert xql.to_dataset(plain, template=template).flag.dtype == bool + assert xql.to_dataset(summed, template=template).flag.dtype == np.int64 + + def test_to_dataset_requires_dims_or_template(): table = pa.table({"a": [1, 2], "b": [3.0, 4.0]}) with pytest.raises(ValueError, match="dims cannot be inferred"): diff --git a/xarray_sql/ds.py b/xarray_sql/ds.py index 6a03deaa..e45b4d23 100644 --- a/xarray_sql/ds.py +++ b/xarray_sql/ds.py @@ -127,19 +127,27 @@ def _as_dtype(coord: xr.DataArray, dtype: np.dtype) -> xr.DataArray: """(result kind, template kind) pairs a database may have widened.""" -def _restore_dtype(var: xr.DataArray, dtype: np.dtype) -> xr.DataArray: - """*var* as the template's *dtype*, when that loses nothing. +def _restore_dtype(var: xr.DataArray, template: xr.DataArray) -> xr.DataArray: + """*var* as the *template* variable's dtype, when *var* is that variable. Databases without a type widen it: SQLite stores ``float32`` as ``float64`` and ``bool`` as an integer, MySQL ``bool`` as ``int8``, and - databases without unsigned integers store them as signed ones. - A plain ``SELECT`` of such a column narrows back exactly; a derived - value (an ``AVG`` of ``float32``, a ``SUM`` of ``bool``) does not, and - keeps the result's dtype. Only in-memory values can be checked, so a - lazily reconstructed variable keeps the result's dtype too. + databases without unsigned integers store them as signed ones. A plain + ``SELECT`` of such a column is narrowed back. + + Whether the column is the variable is read from its dims, since the + query itself is not available here: a plain select (filtered or not) + keeps every dim of the variable, while an aggregate such as + ``SUM(flag) AS flag`` reduces some away and keeps the result's dtype, + whatever its values. The narrowing must also be exact, and only + in-memory values can be checked, so a lazily reconstructed variable + keeps the result's dtype too. """ + dtype = template.dtype if var.dtype == dtype or not isinstance(var.data, np.ndarray): return var + if set(var.dims) != set(template.dims): + return var if (var.dtype.kind, dtype.kind) not in _NARROWINGS: return var values = var.values @@ -182,7 +190,7 @@ def _apply_template(ds: xr.Dataset, template: xr.Dataset) -> xr.Dataset: # template. Aggregation aliases absent from template inherit nothing. for name in list(out.data_vars): if name in template.data_vars: - out[name] = _restore_dtype(out[name], template[name].dtype) + out[name] = _restore_dtype(out[name], template[name]) out[name].attrs = dict(template[name].attrs) # Drop dtype-bound encoding keys; SQL may have changed dtype. enc = { From 95cab62020a5c7eac1f24faf400cabd06e63cc63 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Mon, 28 Sep 2026 08:34:33 -0700 Subject: [PATCH 29/31] Catch only driver errors in the ADBC adapter; replace atomically on ClickHouse - The adapter caught Exception where it probes for driver features and falls back, so a bug or a dropped connection became flat tables and a warning, discovered only when a later query failed. It now catches only the DB-API Error that every ADBC driver raises, looked up from the loaded driver manager so ADBC stays optional. - On ClickHouse, which has no transactions, mode="replace" dropped the table before ingesting, so a failed ingest lost the old data. The ingest now goes into a staging table that is swapped in with EXCHANGE TABLES (atomic) only on success, and dropped on failure. The append-only DDL hook becomes a plan: statements before the ingest, the table it appends to, and statements after success or failure. From review. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- tests/test_adbc_backend.py | 24 ++++++++ xarray_sql/backends/_adbc_dialects.py | 79 +++++++++++++++++++++++---- xarray_sql/backends/adbc.py | 51 +++++++++++------ 3 files changed, 125 insertions(+), 29 deletions(-) diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 8a8d4606..378420fc 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -458,6 +458,30 @@ def test_postgresql_tables_have_statistics_after_register(db, ds): assert cur.fetchone()[0] == 8 * 5 * 6 +def test_clickhouse_failed_replace_keeps_the_old_table(db, ds): + # ClickHouse has no transactions to roll back a half-done replace. + _only(db, "clickhouse", "chdb") + table = db.name("weather") + xql.register(db.con, table, ds) + + def fail_on_second_chunk(block, block_info=None): + if block_info[0]["chunk-location"][0] == 1: + raise OSError("the source went away") + return block + + broken = ds.copy() + broken["temperature"] = broken.temperature.copy( + data=broken.temperature.data.map_blocks( + fail_on_second_chunk, dtype=broken.temperature.dtype + ) + ) + with pytest.raises((OSError, dbapi.Error)): + xql.register(db.con, table, broken, mode="replace") + + out = xql.to_dataset(_select_all(db, table), template=ds) + xr.testing.assert_identical(out, ds.compute()) + + def test_mysql_keeps_the_default_database(db, mixed_ds): # The MySQL driver ignores the target schema, so the adapter switches # the default database for the ingest and must switch it back. diff --git a/xarray_sql/backends/_adbc_dialects.py b/xarray_sql/backends/_adbc_dialects.py index 4118ff1a..ed1d519c 100644 --- a/xarray_sql/backends/_adbc_dialects.py +++ b/xarray_sql/backends/_adbc_dialects.py @@ -15,6 +15,7 @@ from __future__ import annotations import dataclasses +import sys from collections.abc import Callable from typing import TYPE_CHECKING, Literal @@ -26,8 +27,26 @@ IngestMode = Literal["create", "append", "replace", "create_append"] -TableDDL = Callable[..., list[str]] -"""Builds the statements that create a table before an append-only ingest.""" + +@dataclasses.dataclass(frozen=True) +class IngestPlan: + """The DDL around an ingest, for drivers that can only append.""" + + table: str + """The table the rows are appended to.""" + + before: list[str] + """Run before the ingest.""" + + after: list[str] = dataclasses.field(default_factory=list) + """Run once the ingest succeeds.""" + + on_failure: list[str] = dataclasses.field(default_factory=list) + """Run if the ingest fails, before the error propagates.""" + + +TableDDL = Callable[..., IngestPlan] +"""Plans the DDL around an append-only ingest into a table.""" @dataclasses.dataclass(frozen=True) @@ -181,20 +200,28 @@ def _clickhouse_ddl( mode: IngestMode, temporary: bool, database: str | None, -) -> list[str]: - """Statements that prepare *table* for an append-mode ingest. +) -> IngestPlan: + """The DDL around an append-mode ingest into *table*. ClickHouse's ADBC driver only appends, so the table is created here for every other mode. Tables are sorted by their dimensions, so ClickHouse's primary index skips data on dimension predicates the way chunk pruning does in the other engines. + + ClickHouse has no transactions, so ``replace`` ingests into a staging + table and swaps it in with ``EXCHANGE TABLES`` (atomic) only once the + ingest has succeeded: a failed ingest leaves the old table as it was. + Temporary tables cannot be exchanged and are dropped and recreated. """ - if mode == "append": - return [] quote = CLICKHOUSE.quote_identifier - target = quote(table) - if database is not None: - target = f"{quote(database)}.{target}" + + def qualified(name: str) -> str: + if database is None: + return quote(name) + return f"{quote(database)}.{quote(name)}" + + if mode == "append": + return IngestPlan(table, []) columns = ", ".join( f"{quote(field.name)} {_clickhouse_type(field, field.name in dims)}" for field in schema @@ -205,12 +232,29 @@ def _clickhouse_ddl( else: order = ", ".join(quote(dim) for dim in dims) or "tuple()" engine = f"ENGINE = MergeTree ORDER BY ({order})" + target = qualified(table) + if mode == "replace" and not temporary: + staging_name = f"{table}__xarray_sql_replace" + staging = qualified(staging_name) + return IngestPlan( + staging_name, + before=[ + f"DROP TABLE IF EXISTS {staging}", + f"CREATE TABLE {staging} ({columns}) {engine}", + ], + after=[ + f"CREATE TABLE IF NOT EXISTS {target} AS {staging}", + f"EXCHANGE TABLES {target} AND {staging}", + f"DROP TABLE {staging}", + ], + on_failure=[f"DROP TABLE IF EXISTS {staging}"], + ) statements = [] if mode == "replace": statements.append(f"DROP {kind} IF EXISTS {target}") exists = " IF NOT EXISTS" if mode == "create_append" else "" statements.append(f"CREATE {kind}{exists} {target} ({columns}) {engine}") - return statements + return IngestPlan(table, statements) CLICKHOUSE = Dialect( @@ -277,6 +321,17 @@ def _clickhouse_ddl( """Known databases, keyed by a name their drivers' vendor names contain.""" +def driver_error(con: dbapi.Connection) -> type[Exception]: + """The DB-API ``Error`` every ADBC driver's failures derive from. + + Looked up rather than imported, so ADBC stays an optional dependency: + whenever an ADBC connection exists, its driver manager is loaded. + Catching only this lets bugs and other surprises propagate instead of + being mistaken for an unsupported feature. + """ + return sys.modules["adbc_driver_manager.dbapi"].Error + + def dialect_for(con: dbapi.Connection) -> Dialect: """The [Dialect][xarray_sql.backends._adbc_dialects.Dialect] of *con*. @@ -286,7 +341,7 @@ def dialect_for(con: dbapi.Connection) -> Dialect: """ try: info = con.adbc_get_info() - except Exception: # noqa: BLE001 — unimplemented; probe instead + except driver_error(con): # unimplemented; probe instead pass else: vendor = str(info.get("vendor_name") or "").lower() @@ -298,7 +353,7 @@ def dialect_for(con: dbapi.Connection) -> Dialect: with con.cursor() as cur: cur.execute("SELECT 1 FROM system.one") cur.fetchall() - except Exception: # noqa: BLE001 + except driver_error(con): return Dialect("the database") return CLICKHOUSE diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index ffdae3cf..4a8c56bc 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -46,6 +46,7 @@ Dialect, IngestMode, dialect_for, + driver_error, ingestable, ) from .base import register_adapter @@ -80,17 +81,17 @@ def _ingest( """ reader = XarrayPushdownDataset(ds, chunks, **kwargs).scanner().to_reader() reader = ingestable(reader, dialect) + plan = None if dialect.table_ddl is not None: - for statement in dialect.table_ddl( + plan = dialect.table_ddl( table, reader.schema, dims, mode=mode, temporary=temporary, database=db_schema_name, - ): - with con.cursor() as cur: - cur.execute(statement) + ) + _execute(con, plan.before) mode, temporary = "append", False target_schema = db_schema_name in_database: contextlib.AbstractContextManager[None] = ( @@ -102,16 +103,26 @@ def _ingest( ): in_database = _default_database(con, db_schema_name, dialect) target_schema = None - with in_database, con.cursor() as cur: - if ingest_options: - cur.adbc_statement.set_options(**ingest_options) - cur.adbc_ingest( - table, - reader, - mode=mode, - db_schema_name=target_schema, - temporary=temporary, - ) + try: + with in_database, con.cursor() as cur: + if ingest_options: + cur.adbc_statement.set_options(**ingest_options) + cur.adbc_ingest( + plan.table if plan is not None else table, + reader, + mode=mode, + db_schema_name=target_schema, + temporary=temporary, + ) + except BaseException: + if plan is not None: + try: + _execute(con, plan.on_failure) + except driver_error(con): + pass # the ingest's own error is the one to report + raise + if plan is not None: + _execute(con, plan.after) if dialect.analyze: target = dialect.quote_identifier(table) if db_schema_name is not None: @@ -120,6 +131,12 @@ def _ingest( cur.execute(f"ANALYZE {target}") +def _execute(con: dbapi.Connection, statements: list[str]) -> None: + for statement in statements: + with con.cursor() as cur: + cur.execute(statement) + + @contextlib.contextmanager def _default_database( con: dbapi.Connection, database: str, dialect: Dialect @@ -149,7 +166,7 @@ def _schema_exists(con: dbapi.Connection, name: str) -> bool: objects = con.adbc_get_objects( depth="db_schemas", db_schema_filter=name ).read_all() - except Exception: # noqa: BLE001 — metadata unsupported; assume not + except driver_error(con): # metadata unsupported; assume not return False return any( schema["db_schema_name"] == name @@ -164,7 +181,7 @@ def _connection_usable(con: dbapi.Connection) -> bool: with con.cursor() as cur: cur.execute("SELECT 1") cur.fetchall() - except Exception: # noqa: BLE001 + except driver_error(con): return False return True @@ -232,7 +249,7 @@ def _create_schema(con: dbapi.Connection, name: str, dialect: Dialect) -> bool: try: with con.cursor() as cur: cur.execute(dialect.create_schema_sql(name)) - except Exception as exc: + except driver_error(con) as exc: if not _connection_usable(con): raise RuntimeError( f"Could not create the {name!r} schema to hold the dimension " From ce3b79ab140b04250cc852097f107a14e26fe9d9 Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Mon, 28 Sep 2026 08:34:33 -0700 Subject: [PATCH 30/31] Pin CI's database images and ADBC drivers; require adbc-driver-manager 1.12 - SQL Server and GizmoSQL ran from floating tags (2022-latest, latest) and dbc installed the newest drivers, so an upstream release could turn CI red with no change here. They are pinned to the digests and versions of the last green run. - The adbc extra allowed adbc-driver-manager 1.0, but only 1.12 is tested, and the adapter relies on adbc_get_objects filters and statement options; the floor (and the test drivers') is now 1.12. - The docs note ClickHouse's atomic replace. From review. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- .github/workflows/adbc-databases.yml | 11 +++++++---- .github/workflows/ci.yml | 4 ++-- docs/engines.md | 7 +++++-- pyproject.toml | 4 ++-- uv.lock | 4 ++-- 5 files changed, 18 insertions(+), 12 deletions(-) diff --git a/.github/workflows/adbc-databases.yml b/.github/workflows/adbc-databases.yml index c27c239f..4570adbd 100644 --- a/.github/workflows/adbc-databases.yml +++ b/.github/workflows/adbc-databases.yml @@ -57,11 +57,12 @@ jobs: port: 8080 uri: http://ci@localhost:8080?catalog=memory&schema=default - backend: mssql - image: mcr.microsoft.com/mssql/server:2022-latest + # Floating tags pinned to the digests of the last green run. + image: mcr.microsoft.com/mssql/server:2022-latest@sha256:4402d880dd4c34bfa7d8705e56a86cd6c88da80a1f6bbbe741f999e76264a090 port: 1433 uri: sqlserver://sa:XqlPassw0rd@localhost:1433?database=master - backend: flightsql - image: gizmodata/gizmosql:latest + image: gizmodata/gizmosql:latest@sha256:7d42a760fe9ba0bf6afb3580a7d962125eb68a3082d19768e13508a4db7c8a96 port: 31337 uri: grpc://localhost:31337 services: @@ -113,8 +114,10 @@ jobs: - name: Install ADBC drivers run: | - uv pip install dbc adbc-driver-postgresql adbc-driver-flightsql - for driver in chdb datafusion mysql clickhouse trino mssql; do + # Pinned, so an upstream release cannot turn this red unannounced. + uv pip install dbc adbc-driver-postgresql==1.12.0 adbc-driver-flightsql==1.12.0 + for driver in chdb=26.7.0 datafusion=0.27.0 mysql=0.6.1 \ + clickhouse=0.1.1 trino=0.5.3 mssql=1.6.2; do uv run --no-project dbc install "$driver" done diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index eb57c0f4..4a64cab7 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -76,8 +76,8 @@ jobs: VIRTUAL_ENV: ${{ github.workspace }}/.venv run: | uv pip install dbc - uv run --no-project dbc install chdb - uv run --no-project dbc install datafusion + uv run --no-project dbc install "chdb=26.7.0" + uv run --no-project dbc install "datafusion=0.27.0" - name: Run unit tests env: VIRTUAL_ENV: ${{ github.workspace }}/.venv diff --git a/docs/engines.md b/docs/engines.md index 28e7e4ce..1099b2b6 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -297,7 +297,10 @@ rather than the server's local zone. Mixed-dimension Datasets go into a ClickHouse *database* named after the Dataset (`era5.surface`), and `temporary=True` creates `Memory` tables. To choose the engine or sort key yourself, create the table -first and register with `mode="append"`. +first and register with `mode="append"`. ClickHouse has no +transactions, so `mode="replace"` ingests into a staging table and +swaps it in with `EXCHANGE TABLES` only once the ingest succeeds; a +failed replace leaves the old table as it was. **Tested databases.** Databases differ in how they quote identifiers, whether they have schemas, which types they store, and what their @@ -356,7 +359,7 @@ What each integration provides. Known issues and constraints live on | `geometry` column ([geospatial](geospatial.md#geoarrow-point-geometry-columns)) | annotated WKB passes through | native `GEOMETRY` (`"wkb"` encoding) | plain binary/struct | driver-dependent | | Mixed-dimension datasets | one schema, `name.group` tables | `name.group` views over `name_group` tables | `xql.arrow_datasets(ds, name)`, one per group | `name.group` tables in a schema; `name_group` without schemas | | Naming those tables (`table_names=`) | yes | yes | yes | yes | -| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.0` (tested on 1.12; see [tested databases](#adbc-adapter-any-database-with-a-driver)) | +| Version floor | bundled (core dependency) | `duckdb >= 1.4` (tested on 1.5) | tested on `polars 1.42` | `adbc-driver-manager >= 1.12` (see [tested databases](#adbc-adapter-any-database-with-a-driver)) | [^spill-only]: Why DuckDB relations do not re-execute — and two other engine-specific issues worth knowing — is explained on diff --git a/pyproject.toml b/pyproject.toml index b11f02b5..ce5545e8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,7 +38,7 @@ dependencies = [ [project.optional-dependencies] adbc = [ # Plus a driver package for your database, e.g. adbc-driver-postgresql. - "adbc-driver-manager>=1.0", + "adbc-driver-manager>=1.12", ] duckdb = [ "duckdb>=1.4.0", @@ -53,7 +53,7 @@ geo = [ "pyproj", ] test = [ - "adbc-driver-sqlite>=1.0", + "adbc-driver-sqlite>=1.12", "cftime", "xarray-sql[adbc,duckdb,polars,geo]", "pytest", diff --git a/uv.lock b/uv.lock index 97e1aa39..6e2321a4 100644 --- a/uv.lock +++ b/uv.lock @@ -2872,8 +2872,8 @@ dev = [ [package.metadata] requires-dist = [ - { name = "adbc-driver-manager", marker = "extra == 'adbc'", specifier = ">=1.0" }, - { name = "adbc-driver-sqlite", marker = "extra == 'test'", specifier = ">=1.0" }, + { name = "adbc-driver-manager", marker = "extra == 'adbc'", specifier = ">=1.12" }, + { name = "adbc-driver-sqlite", marker = "extra == 'test'", specifier = ">=1.12" }, { name = "cftime", marker = "extra == 'test'" }, { name = "dask", specifier = ">=2024.8.0" }, { name = "datafusion", specifier = "==54.0.0" }, From a48896ffaf9c38e1a290f243f1db9320f4ab919c Mon Sep 17 00:00:00 2001 From: Alex Merose Date: Mon, 28 Sep 2026 08:42:09 -0700 Subject: [PATCH 31/31] Type driver_error's lookup (mypy) Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01AMTHTEAoyKzUJLvFmg5G6t --- xarray_sql/backends/_adbc_dialects.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/xarray_sql/backends/_adbc_dialects.py b/xarray_sql/backends/_adbc_dialects.py index ed1d519c..9e337e2a 100644 --- a/xarray_sql/backends/_adbc_dialects.py +++ b/xarray_sql/backends/_adbc_dialects.py @@ -329,7 +329,8 @@ def driver_error(con: dbapi.Connection) -> type[Exception]: Catching only this lets bugs and other surprises propagate instead of being mistaken for an unsupported feature. """ - return sys.modules["adbc_driver_manager.dbapi"].Error + error: type[Exception] = sys.modules["adbc_driver_manager.dbapi"].Error + return error def dialect_for(con: dbapi.Connection) -> Dialect: