diff --git a/.github/workflows/adbc-databases.yml b/.github/workflows/adbc-databases.yml new file mode 100644 index 00000000..4570adbd --- /dev/null +++ b/.github/workflows/adbc-databases.yml @@ -0,0 +1,162 @@ +# 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: + 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/_adbc.py" + - "tests/test_adbc_backend.py" + - "tests/test_adbc_era5_integration.py" + - ".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: ${{ 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 + # 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@sha256:7d42a760fe9ba0bf6afb3580a7d962125eb68a3082d19768e13508a4db7c8a96 + port: 31337 + uri: grpc://localhost:31337 + services: + # 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 + MYSQL_ROOT_PASSWORD: xql + MYSQL_DATABASE: xql + MARIADB_ROOT_PASSWORD: xql + MARIADB_DATABASE: xql + CLICKHOUSE_PASSWORD: xql + ACCEPT_EULA: "Y" + MSSQL_SA_PASSWORD: XqlPassw0rd + TLS_ENABLED: "0" + GIZMOSQL_USERNAME: xql + GIZMOSQL_PASSWORD: xql + env: + VIRTUAL_ENV: ${{ github.workspace }}/.venv + XARRAY_SQL_TEST_ONLY: ${{ matrix.backend }} + XARRAY_SQL_TEST_FLIGHTSQL_USERNAME: xql + XARRAY_SQL_TEST_FLIGHTSQL_PASSWORD: xql + 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: | + # 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 + + - name: Point the tests' backend table at the server + if: matrix.uri + run: | + 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 + # 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 diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c1d892d9..4a64cab7 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 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=26.7.0" + uv run --no-project dbc install "datafusion=0.27.0" - name: Run unit tests + env: + VIRTUAL_ENV: ${{ github.workspace }}/.venv run: uv run --no-project pytest -v . -m "not integration" 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..1099b2b6 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -190,22 +190,176 @@ 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. 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"}`. +- **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 +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`. + +**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 + +# 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(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"`. 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 +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 (`2021-01-01 04:00:00…`), so plain literals compare correctly; timedeltas and unsigned integers as integers | +| DuckDB | schema | yes | | +| 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 | 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' +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 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` +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`. + ## 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.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 d7bc0d18..54473e22 100644 --- a/docs/limitations.md +++ b/docs/limitations.md @@ -82,6 +82,57 @@ 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. + + **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. + + **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. + ## 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..ce5545e8 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.12", +] duckdb = [ "duckdb>=1.4.0", ] @@ -49,8 +53,9 @@ geo = [ "pyproj", ] test = [ + "adbc-driver-sqlite>=1.12", "cftime", - "xarray-sql[duckdb,polars,geo]", + "xarray-sql[adbc,duckdb,polars,geo]", "pytest", "xarray[io]", "gcsfs", diff --git a/tests/_adbc.py b/tests/_adbc.py new file mode 100644 index 00000000..dd866015 --- /dev/null +++ b/tests/_adbc.py @@ -0,0 +1,222 @@ +"""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). + +``XARRAY_SQL_TEST_ONLY`` (comma-separated names) restricts a run to +those backends, e.g. one CI job per database. +""" + +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") + 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) + 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 new file mode 100644 index 00000000..378420fc --- /dev/null +++ b/tests/test_adbc_backend.py @@ -0,0 +1,493 @@ +"""Tests for the ADBC engine adapter, run against every available database. + +``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 (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 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") + + +NAMES = { + ("time", "lat", "lon"): "surface", + ("time", "level", "lat", "lon"): "atmosphere", +} + + +@pytest.fixture +def ds() -> xr.Dataset: + rng = np.random.default_rng(3) + weather = xr.Dataset( + data_vars=dict( + 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"), + lat=np.linspace(-10.0, 10.0, 5), + lon=np.linspace(0.0, 40.0, 6), + ), + attrs=dict(description="Synthetic weather."), + ).chunk({"time": 4}) + weather["temperature"][0, 0, 0] = np.nan + weather["sst"][0, 0, 1] = np.nan + return weather + + +@pytest.fixture +def mixed_ds() -> xr.Dataset: + rng = np.random.default_rng(11) + return xr.Dataset( + { + "t2m": (["time", "lat", "lon"], rng.random((6, 3, 4))), + "temperature": ( + ["time", "level", "lat", "lon"], + rng.random((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}) + + +@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 _select_all(db, table: str): + return db.query( + f"SELECT time, lat, lon, temperature, count, sst, land FROM {table} " + "ORDER BY time, lat, lon" + ) + + +# 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_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) + + assert out.temperature.dims == ("lat", "lon") + xr.testing.assert_allclose( + out.temperature, ds.temperature.mean("time").compute() + ) + + +def test_existing_table_is_not_overwritten_by_default(db, ds): + table = db.name("weather") + xql.register(db.con, table, ds) + + with pytest.raises(dbapi.Error): + xql.register(db.con, table, ds) + + +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") + + out = xql.to_dataset(_select_all(db, table), template=ds) + + xr.testing.assert_identical(out, ds.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") + + count = db.query(f"SELECT COUNT(*) FROM {table}").fetchone()[0] + assert count == 2 * 8 * 5 * 6 + + +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 + + xql.register(db.con, table, ds, temporary=True) + + temporary = f"{db.backend.temporary_prefix}{table}" + count = db.query(f"SELECT COUNT(*) FROM {temporary}").fetchone()[0] + assert count == 8 * 5 * 6 + + +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 = db.query( + f"SELECT time, level, lat, lon, temperature FROM {table} " + "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()) + + +@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 = 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 + ) + + xr.testing.assert_allclose(out.compute(), forecast.compute()) + assert out.step.dtype == forecast.step.dtype + + +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) + + assert out.temperature.chunks is not None + xr.testing.assert_allclose( + out.temperature.compute(), ds.temperature.compute() + ) + + +def test_time_filters_select_the_right_rows(db, ds): + # A literal means UTC everywhere, and SQLite's text times compare + # with it correctly, including at an inclusive bound. + table = db.name("weather") + xql.register(db.con, table, ds) + 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 BETWEEN {start} AND {end} ORDER BY time, lat, lon" + ) + out = xql.to_dataset(cur, template=ds) + + expected = ds.temperature.isel(time=slice(4, 7)) + 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( + { + "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 + + +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 --------------------------------------------------- + + +def _only(db, *names: str) -> None: + if db.backend.name not in names: + pytest.skip(f"specific to {', '.join(names)}") + + +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"}, + ) + + +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. + _only(db, "postgresql") + + with pytest.raises(RuntimeError, match="rollback"): + xql.register(db.con, db.name("pg_era5"), mixed_ds, table_names=NAMES) + + +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) + + count = db.query(f"SELECT COUNT(*) FROM {name}.surface").fetchone()[0] + 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_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. + _only(db, "mysql", "mariadb") + before = db.query("SELECT DATABASE()").fetchone()[0] + + xql.register(db.con, db.name("era5"), mixed_ds, table_names=NAMES) + + assert db.query("SELECT DATABASE()").fetchone()[0] == before 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 diff --git a/tests/test_duckdb_backend.py b/tests/test_duckdb_backend.py index 26f9204b..e592cfc5 100644 --- a/tests/test_duckdb_backend.py +++ b/tests/test_duckdb_backend.py @@ -133,6 +133,41 @@ 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_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/tests/test_lazy_roundtrip.py b/tests/test_lazy_roundtrip.py index 185191ef..d8148a58 100644 --- a/tests/test_lazy_roundtrip.py +++ b/tests/test_lazy_roundtrip.py @@ -375,3 +375,47 @@ 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) + + +@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/uv.lock b/uv.lock index cf72eb88..6e2321a4 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.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" }, @@ -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_dialects.py b/xarray_sql/backends/_adbc_dialects.py new file mode 100644 index 00000000..9e337e2a --- /dev/null +++ b/xarray_sql/backends/_adbc_dialects.py @@ -0,0 +1,439 @@ +"""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 +import sys +from collections.abc import Callable +from typing import TYPE_CHECKING, Literal + +import numpy as np +import pyarrow as pa + +if TYPE_CHECKING: + from adbc_driver_manager import dbapi + +IngestMode = Literal["create", "append", "replace", "create_append"] + + +@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) +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.""" + + 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.""" + + 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.""" + + 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", + 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, +) -> 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. + """ + quote = CLICKHOUSE.quote_identifier + + 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 + ) + 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})" + 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 IngestPlan(table, statements) + + +CLICKHOUSE = Dialect( + "ClickHouse", + schema_kind="DATABASE", + durations=False, + table_ddl=_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, + unsigned=False, + timestamps_as_text=True, + ), + "duckdb": Dialect("DuckDB"), + # PostgreSQL wraps a uint64 above the int64 range without an error. + "postgresql": Dialect( + "PostgreSQL", + unsigned=False, + timestamp_unit="us", + folds="lower", + analyze=True, + ), + # 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, folds="lower"), + # Trino's driver ignores temporary=True and creates a permanent table. + "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, + durations=False, + unsigned=False, + ), + "databricks": Dialect("Databricks", quote="`", temporary_tables=False), + "snowflake": Dialect("Snowflake", temporary_tables=False, folds="upper"), +} +"""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. + """ + error: type[Exception] = sys.modules["adbc_driver_manager.dbapi"].Error + return error + + +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 driver_error(con): # 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 driver_error(con): + return Dialect("the database") + return CLICKHOUSE + + +_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. + + 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())) + + +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 every column in a type *dialect*'s database stores. + + 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. + """ + 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 = 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 new file mode 100644 index 00000000..4a8c56bc --- /dev/null +++ b/xarray_sql/backends/adbc.py @@ -0,0 +1,390 @@ +"""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 contextlib +import warnings +from collections.abc import Iterator, Mapping +from typing import TYPE_CHECKING, Any, TypeGuard + +import xarray as xr + +from ..df import ( + Chunks, + TableNames, + group_vars_by_dims, + resolve_table_names, + shared_coord_arrays, +) +from ._adbc_dialects import ( + Dialect, + IngestMode, + dialect_for, + driver_error, + ingestable, +) +from .base import register_adapter +from .pyarrow import XarrayPushdownDataset + +if TYPE_CHECKING: + from adbc_driver_manager import dbapi + +__all__ = ["ADBCAdapter"] + + +def _ingest( + con: dbapi.Connection, + table: str, + ds: xr.Dataset, + chunks: Chunks, + *, + dialect: Dialect, + dims: tuple[str, ...], + mode: IngestMode, + temporary: bool, + ingest_options: Mapping[str, str] | None, + 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() + reader = ingestable(reader, dialect) + plan = None + if dialect.table_ddl is not None: + plan = dialect.table_ddl( + table, + reader.schema, + dims, + mode=mode, + temporary=temporary, + database=db_schema_name, + ) + _execute(con, plan.before) + mode, temporary = "append", False + target_schema = db_schema_name + in_database: contextlib.AbstractContextManager[None] = ( + contextlib.nullcontext() + ) + 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 + 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: + target = f"{dialect.quote_identifier(db_schema_name)}.{target}" + with con.cursor() as cur: + 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 +) -> Iterator[None]: + """Make *database* the connection's default while in the block. + + 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 {dialect.quote_identifier(database)}") + try: + yield + finally: + if previous is not None: + with con.cursor() as cur: + cur.execute(f"USE {dialect.quote_identifier(previous)}") + + +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 driver_error(con): # 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 driver_error(con): + return False + 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( + 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. + + 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 + try: + with con.cursor() as cur: + cur.execute(dialect.create_schema_sql(name)) + 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 " + 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 + return _flat(name, f"could not create the {name!r} schema ({exc})") + 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, + ingest_options: Mapping[str, str] | None = None, + **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') + + 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``. + + 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(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 + 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. 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``, + ``coalesce_rows``) to tune the ingest scan. + """ + groups = group_vars_by_dims(ds) + names = resolve_table_names(ds, table_names, case_insensitive=True) + 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." + ) + _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, + name, + ds, + chunks, + dialect=dialect, + dims=next(iter(groups), ()), + mode=mode, + temporary=temporary, + ingest_options=ingest_options, + **kwargs, + ) + return con + + 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] + _ingest( + con, + group if in_schema else f"{name}_{group}", + ds[var_names], + chunks, + dialect=dialect, + dims=dims, + mode=mode, + temporary=temporary, + ingest_options=ingest_options, + 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 diff --git a/xarray_sql/ds.py b/xarray_sql/ds.py index 94339e0d..e45b4d23 100644 --- a/xarray_sql/ds.py +++ b/xarray_sql/ds.py @@ -56,18 +56,109 @@ # --------------------------------------------------------------------------- -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 = ( + "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) + + +_NARROWINGS = {("i", "b"), ("u", "b"), ("i", "u"), ("f", "f")} +"""(result kind, template kind) pairs a database may have widened.""" + + +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 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 + 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 + ): + return var + return var.copy(data=narrowed) def _apply_template(ds: xr.Dataset, template: xr.Dataset) -> xr.Dataset: @@ -83,6 +174,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 @@ -93,10 +186,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]) out[name].attrs = dict(template[name].attrs) # Drop dtype-bound encoding keys; SQL may have changed dtype. enc = { @@ -114,7 +208,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) @@ -1095,13 +1189,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/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()) diff --git a/xarray_sql/roundtrip.py b/xarray_sql/roundtrip.py index e342828f..a592ccc2 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 @@ -39,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 @@ -346,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 " @@ -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,54 @@ 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 _as_timestamps(array: pa.Array) -> pa.Array: + """Times a database returned as text (SQLite has no time type).""" + # 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( + result: Any, path: str, template: xr.Dataset | None = None +) -> None: + """Write a one-shot Arrow result to Parquet, batch by batch. + + 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: raise TypeError( @@ -490,6 +537,25 @@ def _stream_to_parquet(result: Any, path: str) -> None: "Arrow stream." ) schema, batches = opened + 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 + ) + 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 conversions: + columns = list(batch.columns) + for i, (convert, _) in conversions.items(): + columns[i] = convert(columns[i]) + batch = pa.RecordBatch.from_arrays(columns, schema=schema) writer.write_batch(batch)