diff --git a/.github/workflows/adbc-databases.yml b/.github/workflows/adbc-databases.yml index 4570adbd..214fd9ca 100644 --- a/.github/workflows/adbc-databases.yml +++ b/.github/workflows/adbc-databases.yml @@ -2,8 +2,9 @@ # 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. +# DuckDB, chDB, DataFusion, and xarray-sql's own Flight SQL server); the +# in-process job here adds their ARCO-ERA5 queries. The ClickHouse job +# also has ClickHouse query that Flight SQL server over Arrow Flight. name: adbc databases on: @@ -13,12 +14,15 @@ on: paths: - "xarray_sql/backends/adbc.py" - "xarray_sql/backends/_adbc_dialects.py" + - "xarray_sql/backends/flight.py" + - "src/flight.rs" - "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" + - "tests/test_flight_sql.py" - ".github/workflows/adbc-databases.yml" workflow_dispatch: @@ -35,7 +39,7 @@ jobs: fail-fast: false matrix: include: - - backend: sqlite,duckdb,chdb,datafusion + - backend: sqlite,duckdb,chdb,datafusion,served - backend: postgresql image: postgres:18 port: 5432 @@ -72,6 +76,8 @@ jobs: image: ${{ matrix.image }} ports: - ${{ matrix.port || 1 }}:${{ matrix.port || 1 }} + # Lets ClickHouse reach a Flight SQL server the tests start here. + options: --add-host=host.docker.internal:host-gateway env: POSTGRES_PASSWORD: xql MYSQL_ROOT_PASSWORD: xql @@ -155,6 +161,14 @@ jobs: if: matrix.uri run: uv run --no-project pytest -v -rs tests/test_adbc_backend.py + - name: Query xarray-sql's Flight SQL server from ClickHouse + if: matrix.backend == 'clickhouse' + run: >- + uv run --no-project pytest -v -rs tests/test_flight_sql.py + -k clickhouse + env: + XARRAY_SQL_TEST_CLICKHOUSE_FLIGHT_HOST: host.docker.internal + - name: Run realistic ARCO-ERA5 queries # Reads a regional subset anonymously from the public bucket. run: >- diff --git a/Cargo.lock b/Cargo.lock index fd5c09ce..6a70b0c9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -103,9 +103,9 @@ dependencies = [ [[package]] name = "arrow-arith" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a0ab212d2c1886e802f51c5212d78ebbcbb0bec980fff9dadc1eb8d45cd0b738" +checksum = "0a41203398f0eaa6f7ec8e62c0da742a21abf282c148fc157f6c35c90e29981a" dependencies = [ "arrow-array", "arrow-buffer", @@ -117,9 +117,9 @@ dependencies = [ [[package]] name = "arrow-array" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cfd33d3e92f207444098c75b42de99d329562be0cf686b307b097cc52b4e999e" +checksum = "ae33dad492b7df00a217563a7b0ef2874df68a0deea1b1a3acf628152f7f7a69" dependencies = [ "ahash", "arrow-buffer", @@ -136,9 +136,9 @@ dependencies = [ [[package]] name = "arrow-buffer" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0c6cd424c2693bcdbc150d843dc9d4d137dd2de4782ce6df491ad11a3a0416c0" +checksum = "b9552f96391c005e6ab449fa941420935e7e062489b12b8b1b08879b2163f5b5" dependencies = [ "bytes", "half", @@ -148,9 +148,9 @@ dependencies = [ [[package]] name = "arrow-cast" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4c5aefb56a2c02e9e2b30746241058b85f8983f0fcff2ba0c6d09006e1cded7f" +checksum = "3a8a327c9649f30d8406995f27642b68df354713cca3baaaf100f076f18d5f34" dependencies = [ "arrow-array", "arrow-buffer", @@ -185,9 +185,9 @@ dependencies = [ [[package]] name = "arrow-data" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c88210023a2bfee1896af366309a3028fc3bcbd6515fa29a7990ee1baa08ee0" +checksum = "2b24852db04738907e06c04ea61e42fe7fda962a34513022dc0d0e754fb7976b" dependencies = [ "arrow-buffer", "arrow-schema", @@ -196,11 +196,39 @@ dependencies = [ "num-traits", ] +[[package]] +name = "arrow-flight" +version = "58.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2dbe34824c639e43136af8f106992792ab456540d54b880bc320a3192502d2e" +dependencies = [ + "arrow-arith", + "arrow-array", + "arrow-buffer", + "arrow-cast", + "arrow-data", + "arrow-ipc", + "arrow-ord", + "arrow-row", + "arrow-schema", + "arrow-select", + "arrow-string", + "base64", + "bytes", + "futures", + "once_cell", + "paste", + "prost", + "prost-types", + "tonic", + "tonic-prost", +] + [[package]] name = "arrow-ipc" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "238438f0834483703d88896db6fe5a7138b2230debc31b34c0336c2996e3c64f" +checksum = "29a908a11fcfb3fb2f6730f4ac15e367bc644e419155e96238f68cf3adde572b" dependencies = [ "arrow-array", "arrow-buffer", @@ -239,9 +267,9 @@ dependencies = [ [[package]] name = "arrow-ord" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1bffd8fd2579286a5d63bac898159873e5094a79009940bcb42bbfce4f19f1d0" +checksum = "63a083ec750f5c043f02946b4baf05fcdbb55f4560a3277055caca5cc99f3eb0" dependencies = [ "arrow-array", "arrow-buffer", @@ -264,9 +292,9 @@ dependencies = [ [[package]] name = "arrow-row" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bab5994731204603c73ba69267616c50f80780774c6bb0476f1f830625115e0c" +checksum = "514ba0ef0d4c5896202dae736251ce415abb43a950bed570fb7981b8716c0e4c" dependencies = [ "arrow-array", "arrow-buffer", @@ -277,9 +305,9 @@ dependencies = [ [[package]] name = "arrow-schema" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f633dbfdf39c039ada1bf9e34c694816eb71fbb7dc78f613993b7245e078a1ed" +checksum = "21ca356ad6425cecb6eb7b28e4f659f1ee7880fbb1a16127de7dd62901efee9e" dependencies = [ "bitflags", "serde_core", @@ -288,9 +316,9 @@ dependencies = [ [[package]] name = "arrow-select" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8cd065c54172ac787cf3f2f8d4107e0d3fdc26edba76fdf4f4cc170258942222" +checksum = "c58da39eb3d8350ad4a549e5c2bc49284dac554016c69829310350f1731b0aad" dependencies = [ "ahash", "arrow-array", @@ -302,9 +330,9 @@ dependencies = [ [[package]] name = "arrow-string" -version = "58.3.0" +version = "58.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "29dd7cda3ab9692f43a2e4acc444d760cc17b12bb6d8232ddf64e9bab7c06b42" +checksum = "b6789b388467525e3271326b6b4915666ecfdf5142aef09779445c954b67543c" dependencies = [ "arrow-array", "arrow-buffer", @@ -354,7 +382,7 @@ checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -365,7 +393,7 @@ checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -377,12 +405,61 @@ dependencies = [ "num-traits", ] +[[package]] +name = "atomic-waker" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" + [[package]] name = "autocfg" version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" +[[package]] +name = "axum" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90" +dependencies = [ + "axum-core", + "bytes", + "futures-util", + "http", + "http-body", + "http-body-util", + "itoa", + "matchit", + "memchr", + "mime", + "percent-encoding", + "pin-project-lite", + "serde_core", + "sync_wrapper", + "tower", + "tower-layer", + "tower-service", +] + +[[package]] +name = "axum-core" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "http-body-util", + "mime", + "pin-project-lite", + "sync_wrapper", + "tower-layer", + "tower-service", +] + [[package]] name = "base64" version = "0.22.1" @@ -510,9 +587,9 @@ dependencies = [ [[package]] name = "cfg-if" -version = "1.0.3" +version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2fd1289c04a9ea8cb22300a459a72a385d7c73d3259e2ed7dcb2af674838cfa9" +checksum = "4e7648175b45a9a48536d676f68d918270699102aa8dab5496df06904c914600" [[package]] name = "chrono" @@ -676,9 +753,9 @@ dependencies = [ [[package]] name = "dashmap" -version = "6.1.0" +version = "6.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5041cc499144891f3790297212f32a74fb938e5136a14943f338ef9e0ae276cf" +checksum = "e6361d5c062261c78a176addb82d4c821ae42bed6089de0e12603cd25de2059c" dependencies = [ "cfg-if", "crossbeam-utils", @@ -1200,7 +1277,7 @@ checksum = "587164e03ad68732aa9e7bfe5686e3f25970d4c64fd4bd80790749840892dae5" dependencies = [ "datafusion-doc", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -1446,7 +1523,7 @@ checksum = "97369cbbc041bc366949bc74d34658d6cda5621039731c6310521892a3a20ae0" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -1593,7 +1670,7 @@ checksum = "162ee34ebcb7c64a8abebc059ce0fee27c2262618d7b60ed8faf72fef13c3650" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -1676,6 +1753,25 @@ version = "0.3.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280" +[[package]] +name = "h2" +version = "0.4.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ef8e5e5a340588f4452631496976cf8636d4a7ecf600239fdc27615d2530bc16" +dependencies = [ + "atomic-waker", + "bytes", + "fnv", + "futures-core", + "futures-sink", + "http", + "indexmap", + "slab", + "tokio", + "tokio-util", + "tracing", +] + [[package]] name = "half" version = "2.7.1" @@ -1737,6 +1833,41 @@ dependencies = [ "itoa", ] +[[package]] +name = "http-body" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ca2a8f2913ee65f60facd6a5905613afaa448497a0230cc41ce022d93290bc2c" +dependencies = [ + "bytes", + "http", +] + +[[package]] +name = "http-body-util" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23169fe34a5fbcdd3f3862e78fb9b6fccd5f02a6dc6f732547005d45631ce71c" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "pin-project-lite", +] + +[[package]] +name = "httparse" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" + +[[package]] +name = "httpdate" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" + [[package]] name = "humantime" version = "2.3.0" @@ -1752,6 +1883,62 @@ dependencies = [ "typenum", ] +[[package]] +name = "hyper" +version = "1.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "27b501faa50e7a26c3d3560ca625132f4078a17771f4810baf70475ae48cbe43" +dependencies = [ + "atomic-waker", + "bytes", + "futures-channel", + "futures-core", + "h2", + "http", + "http-body", + "httparse", + "httpdate", + "itoa", + "pin-project-lite", + "smallvec", + "tokio", + "want", +] + +[[package]] +name = "hyper-timeout" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b90d566bffbce6a75bd8b09a05aa8c2cb1fabb6cb348f8840c9e4c90a0d83b0" +dependencies = [ + "hyper", + "hyper-util", + "pin-project-lite", + "tokio", + "tower-service", +] + +[[package]] +name = "hyper-util" +version = "0.1.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddc03d96684f9226b8a787cdb71488417b53ab5ea8fdb1dac946cb9431cc8bff" +dependencies = [ + "bytes", + "futures-channel", + "futures-util", + "http", + "http-body", + "httparse", + "hyper", + "libc", + "pin-project-lite", + "socket2", + "tokio", + "tower-service", + "tracing", +] + [[package]] name = "iana-time-zone" version = "0.1.64" @@ -2076,6 +2263,12 @@ dependencies = [ "twox-hash", ] +[[package]] +name = "matchit" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" + [[package]] name = "md-5" version = "0.11.0" @@ -2092,6 +2285,12 @@ version = "2.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4" +[[package]] +name = "mime" +version = "0.3.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" + [[package]] name = "miniz_oxide" version = "0.8.9" @@ -2102,6 +2301,17 @@ dependencies = [ "simd-adler32", ] +[[package]] +name = "mio" +version = "1.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4b18443e9c262bfe8fa82f51666e2642c53393f7e5c27b3e1aeab922cff5b9d8" +dependencies = [ + "libc", + "wasi 0.11.1+wasi-snapshot-preview1", + "windows-sys 0.61.1", +] + [[package]] name = "num-bigint" version = "0.4.6" @@ -2168,9 +2378,9 @@ dependencies = [ [[package]] name = "once_cell" -version = "1.21.3" +version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" [[package]] name = "ordered-float" @@ -2299,7 +2509,7 @@ checksum = "c96395f0a926bc13b1c17622aaddda1ecb55d49c8f1bf9777e4d877800a43f8b" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -2382,7 +2592,16 @@ dependencies = [ "itertools", "proc-macro2", "quote", - "syn", + "syn 2.0.118", +] + +[[package]] +name = "prost-types" +version = "0.14.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8991c4cbdb8bc5b11f0b074ffe286c30e523de90fee5ba8132f1399f23cb3dd7" +dependencies = [ + "prost", ] [[package]] @@ -2436,7 +2655,7 @@ dependencies = [ "proc-macro2", "pyo3-macros-backend", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -2449,7 +2668,7 @@ dependencies = [ "proc-macro2", "pyo3-build-config", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -2519,7 +2738,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "76009fbe0614077fc1a2ce255e3a1881a2e3a3527097d5dc6d8212c585e7e38b" dependencies = [ "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -2623,9 +2842,9 @@ checksum = "1bc711410fbe7399f390ca1c3b60ad0f53f80e95c5eb935e52268a0e2cd49acc" [[package]] name = "serde" -version = "1.0.227" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "80ece43fc6fbed4eb5392ab50c07334d3e577cbf40997ee896fe7af40bba4245" +checksum = "4148590afebada386688f18773da617792bf2ef03ffc1e4cbd2b1d45b023e0ba" dependencies = [ "serde_core", "serde_derive", @@ -2633,22 +2852,22 @@ dependencies = [ [[package]] name = "serde_core" -version = "1.0.227" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7a576275b607a2c86ea29e410193df32bc680303c82f31e275bbfcafe8b33be5" +checksum = "67dca2c9c51e58a4791a4b1ed58308b39c64224d349a935ab5039aa360942a48" dependencies = [ "serde_derive", ] [[package]] name = "serde_derive" -version = "1.0.227" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "51e694923b8824cf0e9b382adf0f60d4e05f348f357b38833a3fa5ed7c2ede04" +checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -2723,6 +2942,16 @@ version = "1.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1b6b67fb9a61334225b5b790716f609cd58395f895b3fe8b328786812a40bc3b" +[[package]] +name = "socket2" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3d1e2c7f27f8d4cb10542a02c49005dbd6e93095799d6f3be745fae9f8fedd4" +dependencies = [ + "libc", + "windows-sys 0.61.1", +] + [[package]] name = "sqlparser" version = "0.62.0" @@ -2742,7 +2971,7 @@ checksum = "a6dd45d8fc1c79299bfbb7190e42ccbbdf6a5f52e4a6ad98d92357ea965bd289" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -2776,7 +3005,7 @@ dependencies = [ "proc-macro-crate", "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -2815,6 +3044,23 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "syn" +version = "3.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8593e8e72159ed2257d083c7a454a85cbf854f37a0966d8d483aff8c8a3ebcee" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "sync_wrapper" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" + [[package]] name = "synstructure" version = "0.13.2" @@ -2823,7 +3069,7 @@ checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -2847,22 +3093,22 @@ dependencies = [ [[package]] name = "thiserror" -version = "2.0.16" +version = "2.0.21" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3467d614147380f2e4e374161426ff399c91084acd2363eaf549172b3d5e60c0" +checksum = "09e52cb86a36cede5cb101bf8908837b3e4c6e5e59fe7fd85c23fb56200d189e" dependencies = [ "thiserror-impl", ] [[package]] name = "thiserror-impl" -version = "2.0.16" +version = "2.0.21" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6c5e1be1c48b9172ee610da68fd9cd2770e7a4056cb3fc98710ee6906f0c7960" +checksum = "fe5197923287db20a58125f0bc85c062f7f2c892de97b18c356f9efb14b28524" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -2902,8 +3148,12 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" dependencies = [ "bytes", + "libc", + "mio", "pin-project-lite", + "socket2", "tokio-macros", + "windows-sys 0.61.1", ] [[package]] @@ -2914,7 +3164,7 @@ checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -2972,6 +3222,77 @@ dependencies = [ "winnow", ] +[[package]] +name = "tonic" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac2a5518c70fa84342385732db33fb3f44bc4cc748936eb5833d2df34d6445ef" +dependencies = [ + "async-trait", + "axum", + "base64", + "bytes", + "h2", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-timeout", + "hyper-util", + "percent-encoding", + "pin-project", + "socket2", + "sync_wrapper", + "tokio", + "tokio-stream", + "tower", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "tonic-prost" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "50849f68853be452acf590cde0b146665b8d507b3b8af17261df47e02c209ea0" +dependencies = [ + "bytes", + "prost", + "tonic", +] + +[[package]] +name = "tower" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" +dependencies = [ + "futures-core", + "futures-util", + "indexmap", + "pin-project-lite", + "slab", + "sync_wrapper", + "tokio", + "tokio-util", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "tower-layer" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "121c2a6cda46980bb0fcd1647ffaf6cd3fc79a013de288782836f6df9c48780e" + +[[package]] +name = "tower-service" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8df9b6e13f2d32c91b9bd719c00d1958837bc7dec474d94952798cc8e69eeec3" + [[package]] name = "tracing" version = "0.1.41" @@ -2991,7 +3312,7 @@ checksum = "81383ab64e72a7a8b8e13130c49e3dab29def6d0c7d76a03087b3cf71c5c6903" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -3003,6 +3324,12 @@ dependencies = [ "once_cell", ] +[[package]] +name = "try-lock" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" + [[package]] name = "twox-hash" version = "2.1.2" @@ -3078,6 +3405,15 @@ dependencies = [ "winapi-util", ] +[[package]] +name = "want" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bfa7760aed19e106de2c7c0b581b509f2f25d3dacaf737cb82ac61bc6d760b0e" +dependencies = [ + "try-lock", +] + [[package]] name = "wasi" version = "0.11.1+wasi-snapshot-preview1" @@ -3125,7 +3461,7 @@ dependencies = [ "log", "proc-macro2", "quote", - "syn", + "syn 2.0.118", "wasm-bindgen-shared", ] @@ -3160,7 +3496,7 @@ checksum = "9f07d2f20d4da7b26400c9f4a0511e6e0345b040694e8a75bd41d578fa4421d7" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", "wasm-bindgen-backend", "wasm-bindgen-shared", ] @@ -3224,7 +3560,7 @@ checksum = "edb307e42a74fb6de9bf3a02d9712678b22399c87e6fa869d6dfcd8c1b7754e0" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -3235,7 +3571,7 @@ checksum = "c0abd1ddbc6964ac14db11c7213d6532ef34bd9aa042c2e5935f59d7908b46a5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -3370,15 +3706,18 @@ name = "xarray_sql" version = "0.4.0" dependencies = [ "arrow", + "arrow-flight", "async-stream", "async-trait", "datafusion", "datafusion-ffi", "futures", "half", + "prost", "pyo3", "pyo3-build-config", "tokio", + "tonic", ] [[package]] @@ -3401,7 +3740,7 @@ checksum = "38da3c9736e16c5d3c8c597a9aaa5d1fa565d0532ae05e27c24aa62fb32c0ab6" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", "synstructure", ] @@ -3422,7 +3761,7 @@ checksum = "88d2b8d9c68ad2b9e4340d7832716a4d21a22a1154777ad56ea55c51a9cf3831" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -3442,7 +3781,7 @@ checksum = "d71e5d6e06ab090c67b5e44993ec16b72dcbaabc526db883a360057678b48502" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", "synstructure", ] @@ -3476,7 +3815,7 @@ checksum = "5b96237efa0c878c64bd89c436f661be4e46b2f3eff1ebb976f7ef2321d2f58f" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 7f891558..d916716b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -19,18 +19,21 @@ exclude = [ [dependencies] arrow = { version = "58", features = ["pyarrow"] } +arrow-flight = { version = "58", features = ["flight-sql"] } async-stream = "0.3" async-trait = "0.1" datafusion = { version = "54.0.0" } datafusion-ffi = { version = "54.0.0" } futures = { version = "0.3" } half = "2.7" +prost = "0.14" # `abi3-py310` builds against CPython's stable ABI, so a single wheel per # platform works on all CPython >= 3.10 (matching `requires-python`). Maturin # enables `pyo3/extension-module` through pyproject.toml for wheel builds; it # must stay disabled for ordinary Cargo test binaries so they link libpython. pyo3 = { version = "0.28.0", features = ["abi3-py310"] } -tokio = { version = "1.46.1", features = ["rt"] } +tokio = { version = "1.46.1", features = ["macros", "net", "rt", "rt-multi-thread", "sync", "time"] } +tonic = { version = "0.14", features = ["transport"] } [build-dependencies] diff --git a/README.md b/README.md index 2c322950..0fa10528 100644 --- a/README.md +++ b/README.md @@ -80,7 +80,9 @@ 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. +brings results back. Or go the other way without copying anything: +`xql.serve({'air': ds})` starts an Arrow Flight SQL server that any Flight SQL +client (ADBC, JDBC, ODBC) can query, with the same chunk pruning as in process. `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. diff --git a/docs/engines.md b/docs/engines.md index 1099b2b6..cc6e9b26 100644 --- a/docs/engines.md +++ b/docs/engines.md @@ -344,6 +344,130 @@ 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`. +## Serving over Flight SQL (no copy) + +The adapters above bring an engine to the data. `xql.serve` does the +reverse and brings the data to remote clients: it starts an +[Arrow Flight SQL](https://arrow.apache.org/docs/format/FlightSql.html) +server over the same lazy tables `XarrayContext` uses. + +```python +import xarray_sql as xql + +server = xql.serve({"era5": ds}, port=8815) +server.wait() # in a script: block until Ctrl+C +``` + +Any Flight SQL client on the same machine can then query `era5`; to +serve other machines, read [exposing a +server](#things-to-know-before-exposing-a-server) first. Clients include +ADBC's Flight SQL driver (Python, R, Go, Java), the Flight SQL +JDBC and ODBC drivers, and the SQL tools built on them. From Python: + +```sh +pip install adbc-driver-flightsql +``` + +```python +import adbc_driver_flightsql.dbapi as flight_sql + +con = flight_sql.connect("grpc://localhost:8815") +cur = con.cursor() +cur.execute(""" + SELECT time, AVG(t2m) AS t2m FROM era5 + WHERE lat BETWEEN 40 AND 41 + GROUP BY time ORDER BY time +""") +out = xql.to_dataset(cur, template=ds) # template: the same Dataset, opened client-side +``` + +Nothing is copied. Each query is planned by DataFusion on the server +against lazy tables, so partition pruning on dimension predicates and +projection pushdown work exactly as they do in process: only the chunks +and variables a query touches are read from the source, only while the +query runs, and results stream back as Arrow record batches. + +`xql.register(server, name, ds, table_names=...)` works on a +`xql.FlightSQLServer()` like on any other engine, including while it +serves. Mixed-dimension Datasets are served as `name.group` tables, and +clients can list tables with the standard Flight SQL metadata calls +(e.g. ADBC's `adbc_get_objects`). + +### Tested clients + +**Spark** reads through the +[Arrow Flight SQL JDBC driver](https://arrow.apache.org/docs/java/flight_sql_jdbc_driver.html) +and pushes its column selection and filters into the SQL it sends, so +they reach the server's chunk pruning: + +```python +spark = ( + SparkSession.builder + .config("spark.jars", "flight-sql-jdbc-driver-19.0.0.jar") + .config("spark.driver.extraJavaOptions", "-Duser.timezone=UTC") + .config("spark.sql.session.timeZone", "UTC") + .getOrCreate() +) +era5 = ( + spark.read.format("jdbc") + .option("url", "jdbc:arrow-flight-sql://server-host:8815/?useEncryption=false") + .option("driver", "org.apache.arrow.driver.jdbc.ArrowFlightJdbcDriver") + .option("dbtable", "era5") # or "era5.surface", or "(SELECT ...) AS t" + .load() +) +xql.to_dataset(era5.where("lat > 40").toArrow(), template=ds) +``` + +Run Spark's JVM in UTC (`-Duser.timezone=UTC`, as above). The JDBC +path shifts timestamps by the JVM's zone otherwise; the session time +zone alone does not prevent it. + +**ClickHouse** (25.8+) reads through its `arrowFlight` table function, +which speaks plain Arrow Flight rather than Flight SQL. The dataset name +is a table, or a query that runs on the server: + +```sql +SELECT avg(temperature) FROM arrowFlight('server-host:8815', 'era5.surface'); + +-- ClickHouse does not push filters into arrowFlight; put them in the +-- name to get the server's chunk pruning: +SELECT * FROM arrowFlight('server-host:8815', + 'SELECT time, t2m FROM era5.surface WHERE lat BETWEEN 40 AND 41'); +``` + +Times arrive without a zone, so ClickHouse parses literals compared +with them in its server zone. Add `SETTINGS session_timezone = 'UTC'` +to queries that filter on time. + +### Things to know before exposing a server + +- **No authentication or TLS.** The server binds to `127.0.0.1` by + default. To accept remote connections, bind `0.0.0.0` only inside a + trusted network, or put it behind a proxy that authenticates and + terminates TLS: + + ```python + server = xql.serve({"era5": ds}, host="0.0.0.0", port=8815, memory_limit=8 * 2**30) + ``` +- **Read-only SQL.** DDL, DML, and other statements (`CREATE EXTERNAL + TABLE`, `COPY`, `SET`, ...) are rejected, so clients cannot read or + write the server's filesystem. +- **DataFusion's SQL, without xarray-sql's Python UDFs.** The + `cftime()` and `reproject()` functions `XarrayContext` registers are + not available on the server. +- **Bound its memory.** Any client can send an expensive `ORDER BY`, + join, or aggregation. `memory_limit=` (bytes) caps what those hold at + once: a query that needs more spills to disk where it can and fails + otherwise, and the server keeps serving. It is unbounded by default. + Chunk reads aren't counted, so also cap the process itself (a + container or cgroup memory limit) before serving untrusted clients. +- **Only the served Datasets are reachable.** Besides rejecting writes, + the server doesn't resolve file paths or URLs as tables + (`SELECT * FROM '/etc/hosts'` fails). +- **One process serves every query.** Chunk reads happen in the server + process, so size it (and `chunks=`) for the concurrent load you + expect. + ## Engine support matrix What each integration provides. Known issues and constraints live on diff --git a/pyproject.toml b/pyproject.toml index ce5545e8..75365869 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -53,6 +53,7 @@ geo = [ "pyproj", ] test = [ + "adbc-driver-flightsql>=1.12", "adbc-driver-sqlite>=1.12", "cftime", "xarray-sql[adbc,duckdb,polars,geo]", diff --git a/src/flight.rs b/src/flight.rs new file mode 100644 index 00000000..f41e7d10 --- /dev/null +++ b/src/flight.rs @@ -0,0 +1,715 @@ +//! An Arrow Flight SQL endpoint over lazily registered xarray tables. +//! +//! Remote clients (the ADBC Flight SQL driver, the Flight SQL JDBC/ODBC +//! drivers, and anything else that speaks the protocol) send SQL; a +//! native DataFusion session plans it against the same +//! `PrunableStreamingTable` providers the in-process engine uses, so +//! partition pruning, projection pushdown, and exact statistics carry +//! over unchanged. Results stream back as Arrow record batches, and +//! source chunks are read only while a query executes. +//! +//! Plain Arrow Flight clients (e.g. ClickHouse's `arrowFlight` table +//! function) are served too: a path descriptor names a table, as in +//! `["weather"]` or `["era5.surface"]`, or is itself a `SELECT`/`WITH` +//! query, which is how such clients get pruning they cannot push down. +//! +//! The service is stateless: a statement's ticket and a prepared +//! statement's handle are the SQL text itself, re-planned on use. The +//! session is read-only: DDL (which could read or write the server's +//! filesystem through `CREATE EXTERNAL TABLE` / `COPY`), DML, and other +//! statements are rejected before planning. + +use std::net::TcpListener; +use std::pin::Pin; +use std::sync::{Arc, LazyLock}; +use std::thread::JoinHandle; +use std::time::Duration; + +use arrow::datatypes::Schema; +use arrow::ipc::writer::IpcWriteOptions; +use arrow_flight::encode::FlightDataEncoderBuilder; +use arrow_flight::error::FlightError; +use arrow_flight::flight_descriptor::DescriptorType; +use arrow_flight::flight_service_server::{FlightService, FlightServiceServer}; +use arrow_flight::sql::metadata::{SqlInfoData, SqlInfoDataBuilder}; +use arrow_flight::sql::server::FlightSqlService; +use arrow_flight::sql::{ + ActionClosePreparedStatementRequest, ActionCreatePreparedStatementRequest, + ActionCreatePreparedStatementResult, Any, Command, CommandGetCatalogs, CommandGetDbSchemas, + CommandGetSqlInfo, CommandGetTables, CommandPreparedStatementQuery, CommandStatementQuery, + ProstMessageExt, SqlInfo, TicketStatementQuery, +}; +use arrow_flight::{ + Action, Criteria, Empty, FlightData, FlightDescriptor, FlightEndpoint, FlightInfo, + HandshakeRequest, IpcMessage, PollInfo, SchemaAsIpc, SchemaResult, Ticket, +}; +use datafusion::arrow::record_batch::RecordBatch; +use datafusion::catalog::{MemorySchemaProvider, SchemaProvider}; +use datafusion::error::DataFusionError; +use datafusion::execution::context::SQLOptions; +use datafusion::execution::runtime_env::RuntimeEnvBuilder; +use datafusion::prelude::{DataFrame, SessionConfig, SessionContext}; +use datafusion::sql::TableReference; +use futures::{stream, Stream, TryStreamExt}; +use prost::Message; +use pyo3::exceptions::{PyRuntimeError, PyValueError}; +use pyo3::prelude::*; +use tokio::sync::oneshot; +use tonic::transport::server::TcpIncoming; +use tonic::transport::Server; +use tonic::{Request, Response, Status, Streaming}; + +use crate::LazyArrowStreamTable; + +type DoGetStream = Pin> + Send + 'static>>; + +static SQL_INFO: LazyLock = LazyLock::new(|| { + let mut builder = SqlInfoDataBuilder::new(); + builder.append(SqlInfo::FlightSqlServerName, "xarray-sql"); + builder.append(SqlInfo::FlightSqlServerVersion, env!("CARGO_PKG_VERSION")); + builder.append(SqlInfo::FlightSqlServerArrowVersion, "1.3"); + builder.append(SqlInfo::FlightSqlServerReadOnly, true); + builder.build().expect("static SqlInfo is valid") +}); + +fn read_only() -> SQLOptions { + SQLOptions::new() + .with_allow_ddl(false) + .with_allow_dml(false) + .with_allow_statements(false) +} + +fn plan_error(e: DataFusionError) -> Status { + Status::invalid_argument(e.to_string()) +} + +fn internal_error(e: impl std::fmt::Display) -> Status { + Status::internal(e.to_string()) +} + +fn invalid_argument(e: impl std::fmt::Display) -> Status { + Status::invalid_argument(e.to_string()) +} + +fn utf8(bytes: &[u8]) -> Result<&str, Status> { + std::str::from_utf8(bytes).map_err(invalid_argument) +} + +/// Encode one metadata batch (catalog, schema, and table listings) as a +/// Flight data stream. +fn single_batch_stream( + schema: Arc, + batch: Result, +) -> DoGetStream { + let stream = FlightDataEncoderBuilder::new() + .with_schema(schema) + .build(stream::once(async move { batch })) + .map_err(Status::from); + Box::pin(stream) +} + +/// A FlightInfo whose single endpoint redeems ``ticket`` on this server. +fn flight_info( + schema: &Schema, + ticket: impl ProstMessageExt, + descriptor: FlightDescriptor, +) -> Result, Status> { + let ticket = Ticket::new(ticket.as_any().encode_to_vec()); + let info = FlightInfo::new() + .try_with_schema(schema) + .map_err(internal_error)? + .with_endpoint(FlightEndpoint::new().with_ticket(ticket)) + .with_descriptor(descriptor); + Ok(Response::new(info)) +} + +/// The Flight SQL service over one DataFusion session. +struct XarrayFlightSql { + ctx: SessionContext, +} + +impl XarrayFlightSql { + async fn plan(&self, sql: &str) -> Result { + self.ctx + .sql_with_options(sql, read_only()) + .await + .map_err(plan_error) + } + + async fn result_schema(&self, sql: &str) -> Result, Status> { + Ok(Arc::clone(self.plan(sql).await?.schema().inner())) + } + + async fn execute(&self, sql: &str) -> Result, Status> { + let frame = self.plan(sql).await?; + let schema = Arc::clone(frame.schema().inner()); + let batches = frame.execute_stream().await.map_err(internal_error)?; + let stream = FlightDataEncoderBuilder::new() + .with_schema(schema) + .build(batches.map_err(|e| FlightError::ExternalError(Box::new(e)))) + .map_err(Status::from); + Ok(Response::new(Box::pin(stream))) + } + + /// Every (catalog, schema) pair in the session. + fn schemas(&self) -> Vec<(String, String, Arc)> { + let mut out = Vec::new(); + for catalog_name in self.ctx.catalog_names() { + let Some(catalog) = self.ctx.catalog(&catalog_name) else { + continue; + }; + for schema_name in catalog.schema_names() { + if let Some(schema) = catalog.schema(&schema_name) { + out.push((catalog_name.clone(), schema_name, schema)); + } + } + } + out + } +} + +#[tonic::async_trait] +impl FlightSqlService for XarrayFlightSql { + type FlightService = XarrayFlightSql; + + async fn get_flight_info_statement( + &self, + query: CommandStatementQuery, + request: Request, + ) -> Result, Status> { + let schema = self.result_schema(&query.query).await?; + let ticket = TicketStatementQuery { + statement_handle: query.query.into_bytes().into(), + }; + flight_info(&schema, ticket, request.into_inner()) + } + + async fn do_get_statement( + &self, + ticket: TicketStatementQuery, + _request: Request, + ) -> Result, Status> { + self.execute(utf8(&ticket.statement_handle)?).await + } + + async fn do_action_create_prepared_statement( + &self, + query: ActionCreatePreparedStatementRequest, + _request: Request, + ) -> Result { + let schema = self.result_schema(&query.query).await?; + let IpcMessage(dataset_schema) = SchemaAsIpc::new(&schema, &IpcWriteOptions::default()) + .try_into() + .map_err(internal_error)?; + Ok(ActionCreatePreparedStatementResult { + prepared_statement_handle: query.query.into_bytes().into(), + dataset_schema, + parameter_schema: Default::default(), + }) + } + + async fn do_action_close_prepared_statement( + &self, + _query: ActionClosePreparedStatementRequest, + _request: Request, + ) -> Result<(), Status> { + // Handles are the SQL text; there is nothing to release. + Ok(()) + } + + async fn get_flight_info_prepared_statement( + &self, + query: CommandPreparedStatementQuery, + request: Request, + ) -> Result, Status> { + let schema = self + .result_schema(utf8(&query.prepared_statement_handle)?) + .await?; + flight_info(&schema, query, request.into_inner()) + } + + async fn do_get_prepared_statement( + &self, + query: CommandPreparedStatementQuery, + _request: Request, + ) -> Result, Status> { + self.execute(utf8(&query.prepared_statement_handle)?).await + } + + async fn get_flight_info_catalogs( + &self, + query: CommandGetCatalogs, + request: Request, + ) -> Result, Status> { + let schema = query.into_builder().schema(); + flight_info(&schema, query, request.into_inner()) + } + + async fn do_get_catalogs( + &self, + query: CommandGetCatalogs, + _request: Request, + ) -> Result, Status> { + let mut builder = query.into_builder(); + for catalog in self.ctx.catalog_names() { + builder.append(catalog); + } + Ok(Response::new(single_batch_stream( + builder.schema(), + builder.build(), + ))) + } + + async fn get_flight_info_schemas( + &self, + query: CommandGetDbSchemas, + request: Request, + ) -> Result, Status> { + let schema = query.clone().into_builder().schema(); + flight_info(&schema, query, request.into_inner()) + } + + async fn do_get_schemas( + &self, + query: CommandGetDbSchemas, + _request: Request, + ) -> Result, Status> { + let mut builder = query.into_builder(); + for (catalog, schema, _) in self.schemas() { + builder.append(catalog, schema); + } + Ok(Response::new(single_batch_stream( + builder.schema(), + builder.build(), + ))) + } + + async fn get_flight_info_tables( + &self, + query: CommandGetTables, + request: Request, + ) -> Result, Status> { + let schema = query.clone().into_builder().schema(); + flight_info(&schema, query, request.into_inner()) + } + + async fn do_get_tables( + &self, + query: CommandGetTables, + _request: Request, + ) -> Result, Status> { + let mut builder = query.into_builder(); + for (catalog, schema_name, schema) in self.schemas() { + for table_name in schema.table_names() { + let Some(table) = schema.table(&table_name).await.map_err(internal_error)? else { + continue; + }; + builder + .append( + &catalog, + &schema_name, + &table_name, + "TABLE", + &table.schema(), + ) + .map_err(internal_error)?; + } + } + Ok(Response::new(single_batch_stream( + builder.schema(), + builder.build(), + ))) + } + + async fn get_flight_info_sql_info( + &self, + query: CommandGetSqlInfo, + request: Request, + ) -> Result, Status> { + let schema = query.clone().into_builder(&SQL_INFO).schema(); + flight_info(&schema, query, request.into_inner()) + } + + async fn do_get_sql_info( + &self, + query: CommandGetSqlInfo, + _request: Request, + ) -> Result, Status> { + let builder = query.into_builder(&SQL_INFO); + Ok(Response::new(single_batch_stream( + builder.schema(), + builder.build(), + ))) + } + + async fn register_sql_info(&self, _id: i32, _result: &SqlInfo) {} +} + +/// Serves plain Arrow Flight path descriptors alongside Flight SQL. +/// +/// arrow-flight's `FlightService` implementation for a `FlightSqlService` +/// decodes every descriptor as a Flight SQL command and leaves +/// `GetSchema` unimplemented. This wrapper answers path descriptors and +/// `GetSchema` itself and hands everything else to the Flight SQL +/// service. A path's ticket is an ordinary statement ticket, so `DoGet` +/// needs no special handling. +struct FlightRouter { + sql: XarrayFlightSql, +} + +impl FlightRouter { + /// The SQL a path descriptor stands for, or `None` for a command. + fn path_query(&self, descriptor: &FlightDescriptor) -> Result, Status> { + if descriptor.r#type() != DescriptorType::Path { + return Ok(None); + } + let reference = match descriptor.path.as_slice() { + [name] => { + let first = name.split_whitespace().next().unwrap_or_default(); + if first.eq_ignore_ascii_case("select") || first.eq_ignore_ascii_case("with") { + return Ok(Some(name.clone())); + } + // A registered name matches exactly, as the two- and + // three-part forms do. Otherwise it is parsed like a + // table name in SQL: `era5.surface` is schema-qualified. + let exact = TableReference::bare(name.as_str()); + if self.sql.ctx.table_exist(exact.clone()).unwrap_or(false) { + exact + } else { + TableReference::parse_str(name) + } + } + [schema, table] => TableReference::partial(schema.as_str(), table.as_str()), + [catalog, schema, table] => { + TableReference::full(catalog.as_str(), schema.as_str(), table.as_str()) + } + _ => { + return Err(Status::invalid_argument( + "a Flight descriptor path is a table name or a SQL query", + )) + } + }; + Ok(Some(format!( + "SELECT * FROM {}", + reference.to_quoted_string() + ))) + } + + /// The SQL behind a descriptor: a path, or a Flight SQL statement. + fn descriptor_query(&self, descriptor: &FlightDescriptor) -> Result { + if let Some(sql) = self.path_query(descriptor)? { + return Ok(sql); + } + // A descriptor that doesn't decode is the client's error. + let message = Any::decode(&*descriptor.cmd).map_err(invalid_argument)?; + match Command::try_from(message).map_err(invalid_argument)? { + Command::CommandStatementQuery(query) => Ok(query.query), + Command::CommandPreparedStatementQuery(query) => { + Ok(utf8(&query.prepared_statement_handle)?.to_string()) + } + other => Err(Status::unimplemented(format!( + "GetSchema is not supported for {}", + other.type_url() + ))), + } + } +} + +#[tonic::async_trait] +impl FlightService for FlightRouter { + type HandshakeStream = ::HandshakeStream; + type ListFlightsStream = ::ListFlightsStream; + type DoGetStream = ::DoGetStream; + type DoPutStream = ::DoPutStream; + type DoExchangeStream = ::DoExchangeStream; + type DoActionStream = ::DoActionStream; + type ListActionsStream = ::ListActionsStream; + + async fn get_flight_info( + &self, + request: Request, + ) -> Result, Status> { + let Some(sql) = self.path_query(request.get_ref())? else { + return FlightService::get_flight_info(&self.sql, request).await; + }; + let schema = self.sql.result_schema(&sql).await?; + let ticket = TicketStatementQuery { + statement_handle: sql.into_bytes().into(), + }; + flight_info(&schema, ticket, request.into_inner()) + } + + async fn get_schema( + &self, + request: Request, + ) -> Result, Status> { + let sql = self.descriptor_query(request.get_ref())?; + let schema = self.sql.result_schema(&sql).await?; + let result = SchemaAsIpc::new(&schema, &IpcWriteOptions::default()) + .try_into() + .map_err(internal_error)?; + Ok(Response::new(result)) + } + + async fn handshake( + &self, + request: Request>, + ) -> Result, Status> { + FlightService::handshake(&self.sql, request).await + } + + async fn list_flights( + &self, + request: Request, + ) -> Result, Status> { + FlightService::list_flights(&self.sql, request).await + } + + async fn poll_flight_info( + &self, + request: Request, + ) -> Result, Status> { + FlightService::poll_flight_info(&self.sql, request).await + } + + async fn do_get( + &self, + request: Request, + ) -> Result, Status> { + FlightService::do_get(&self.sql, request).await + } + + async fn do_put( + &self, + request: Request>, + ) -> Result, Status> { + FlightService::do_put(&self.sql, request).await + } + + async fn do_exchange( + &self, + request: Request>, + ) -> Result, Status> { + FlightService::do_exchange(&self.sql, request).await + } + + async fn do_action( + &self, + request: Request, + ) -> Result, Status> { + FlightService::do_action(&self.sql, request).await + } + + async fn list_actions( + &self, + request: Request, + ) -> Result, Status> { + FlightService::list_actions(&self.sql, request).await + } +} + +/// How long in-flight queries may run once a dropped server stops. +const DROP_GRACE_PERIOD: Duration = Duration::from_secs(5); + +/// Handles of a server that is accepting connections. +struct Running { + /// Starts shutdown; carries how long in-flight queries may run on. + shutdown: oneshot::Sender, + /// Ends with the error that stopped the server, if one did. + thread: JoinHandle>, +} + +/// A Flight SQL server over a native DataFusion session. +/// +/// Register tables first, then call ``serve``. The server runs on its own +/// thread with a multi-threaded Tokio runtime; partitions acquire the GIL +/// only for each Python call, as they do in-process. +#[pyclass(name = "FlightSqlServer")] +pub(crate) struct FlightSqlServer { + ctx: SessionContext, + running: Option, +} + +fn runtime_error(e: impl std::fmt::Display) -> PyErr { + PyRuntimeError::new_err(e.to_string()) +} + +#[pymethods] +impl FlightSqlServer { + /// ``memory_limit`` caps, in bytes, the memory queries' sorts, joins, + /// and aggregations may hold; a query that needs more spills to disk + /// where it can and fails otherwise. ``None`` leaves it unbounded. + #[new] + #[pyo3(signature = (memory_limit=None))] + fn new(memory_limit: Option) -> PyResult { + let mut runtime = RuntimeEnvBuilder::new(); + if let Some(limit) = memory_limit { + runtime = runtime.with_memory_limit(limit, 1.0); + } + let runtime = runtime.build_arc().map_err(runtime_error)?; + Ok(Self { + ctx: SessionContext::new_with_config_rt(SessionConfig::new(), runtime), + running: None, + }) + } + + /// Register ``table`` as ``name``. + fn register_table(&self, name: &str, table: PyRef<'_, LazyArrowStreamTable>) -> PyResult<()> { + // Bare, not parsed as SQL: a `&str` would fold `Weather` to + // `weather`, and a quoted `"Weather"` could never find it. + self.ctx + .register_table(TableReference::bare(name), table.table.clone()) + .map_err(runtime_error)?; + Ok(()) + } + + /// Register ``tables`` (name, table) in the SQL schema ``schema``, all + /// or none: a name already taken fails before anything is added. + fn register_tables( + &self, + schema: &str, + tables: Vec<(String, PyRef<'_, LazyArrowStreamTable>)>, + ) -> PyResult<()> { + let catalog = self + .ctx + .catalog("datafusion") + .ok_or_else(|| runtime_error("the default catalog is missing"))?; + let existing = catalog.schema(schema); + if let Some(existing) = &existing { + let taken: Vec<&str> = tables + .iter() + .map(|(name, _)| name.as_str()) + .filter(|name| existing.table_exist(name)) + .collect(); + if !taken.is_empty() { + return Err(runtime_error(format!( + "{schema} already has tables named {}", + taken.join(", ") + ))); + } + } + // A new schema is filled before it is added, so queries never see + // it partly registered. + let is_new = existing.is_none(); + let target: Arc = + existing.unwrap_or_else(|| Arc::new(MemorySchemaProvider::new())); + for (name, table) in &tables { + target + .register_table(name.clone(), table.table.clone()) + .map_err(runtime_error)?; + } + if is_new { + catalog + .register_schema(schema, target) + .map_err(runtime_error)?; + } + Ok(()) + } + + /// Start accepting connections on ``host:port``; returns the bound + /// port (useful with ``port=0``, which picks a free one). + fn serve(&mut self, host: &str, port: u16) -> PyResult { + if self.is_running() { + return Err(runtime_error("the server is already running")); + } + let listener = TcpListener::bind((host, port))?; + listener.set_nonblocking(true)?; + let bound = listener.local_addr()?.port(); + + let service = XarrayFlightSql { + ctx: self.ctx.clone(), + }; + let (shutdown, shutdown_rx) = oneshot::channel::(); + let (ready, ready_rx) = std::sync::mpsc::channel::>(); + let thread = std::thread::Builder::new() + .name("xarray-sql-flight-sql".to_string()) + .spawn(move || { + let runtime = match tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build() + { + Ok(runtime) => runtime, + Err(e) => { + let _ = ready.send(Err(e.to_string())); + return Ok(()); + } + }; + let result = runtime.block_on(async move { + let listener = match tokio::net::TcpListener::from_std(listener) { + Ok(listener) => listener, + Err(e) => { + let _ = ready.send(Err(e.to_string())); + return Ok(()); + } + }; + let incoming = TcpIncoming::from(listener).with_nodelay(Some(true)); + let _ = ready.send(Ok(())); + // Graceful shutdown waits for every open response + // stream, including ones a client stopped reading, so + // it gets a deadline after which the server is dropped. + let (grace, grace_rx) = oneshot::channel::(); + let server = Server::builder() + .add_service(FlightServiceServer::new(FlightRouter { sql: service })) + .serve_with_incoming_shutdown(incoming, async move { + let period = shutdown_rx.await.unwrap_or(Duration::ZERO); + let _ = grace.send(period); + }); + let deadline = async move { + match grace_rx.await { + Ok(period) => tokio::time::sleep(period).await, + Err(_) => std::future::pending::<()>().await, + } + }; + tokio::select! { + served = server => served.map_err(|e| e.to_string()), + _ = deadline => Ok(()), + } + }); + // Cancels whatever the deadline cut off, closing its + // connections, without waiting on it. + runtime.shutdown_background(); + result + })?; + + match ready_rx.recv() { + Ok(Ok(())) => {} + Ok(Err(e)) => return Err(runtime_error(e)), + Err(_) => return Err(runtime_error("the server thread exited during startup")), + } + self.running = Some(Running { shutdown, thread }); + Ok(bound) + } + + /// Whether the server thread is still accepting connections. + fn is_running(&self) -> bool { + self.running + .as_ref() + .is_some_and(|running| !running.thread.is_finished()) + } + + /// Stop accepting connections, give in-flight queries up to + /// ``timeout`` seconds to finish, then close the remaining connections. + /// Raises the error that stopped the server, if one did. + #[pyo3(signature = (timeout=5.0))] + fn shutdown(&mut self, py: Python<'_>, timeout: f64) -> PyResult<()> { + let grace = Duration::try_from_secs_f64(timeout) + .map_err(|e| PyValueError::new_err(format!("invalid timeout {timeout}: {e}")))?; + if let Some(running) = self.running.take() { + let _ = running.shutdown.send(grace); + // In-flight partitions may need the GIL to finish. + py.detach(|| running.thread.join()) + .map_err(|_| runtime_error("the server thread panicked"))? + .map_err(|e| runtime_error(format!("the server stopped: {e}")))?; + } + Ok(()) + } +} + +impl Drop for FlightSqlServer { + fn drop(&mut self) { + // Signal only: joining here could wait on the GIL this thread holds. + if let Some(running) = self.running.take() { + let _ = running.shutdown.send(DROP_GRACE_PERIOD); + } + } +} diff --git a/src/lib.rs b/src/lib.rs index 1111a1b2..61ab2c03 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -41,6 +41,8 @@ //! Will skip loading partitions whose time ranges are entirely before 2020-02-01. //! Supported operators: `=`, `<`, `>`, `<=`, `>=`, `BETWEEN`, `IN`, `AND`, `OR`. +mod flight; + use std::collections::{HashMap, HashSet}; use std::ffi::CString; use std::fmt::Debug; @@ -1522,9 +1524,9 @@ fn ffi_logical_codec_from_pycapsule( /// ``` #[pyclass(name = "LazyArrowStreamTable")] -struct LazyArrowStreamTable { +pub(crate) struct LazyArrowStreamTable { /// The underlying table provider with pruning support - table: Arc, + pub(crate) table: Arc, } #[pymethods] @@ -1643,5 +1645,6 @@ impl LazyArrowStreamTable { #[pymodule] fn _native(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; + m.add_class::()?; Ok(()) } diff --git a/tests/_adbc.py b/tests/_adbc.py index dd866015..4d7d95ae 100644 --- a/tests/_adbc.py +++ b/tests/_adbc.py @@ -13,6 +13,13 @@ ``dbc`` did not install; ``XARRAY_SQL_TEST_FLIGHTSQL_USERNAME`` and ``_PASSWORD`` for a Flight SQL server such as GizmoSQL). +``served`` is xarray-sql's own Flight SQL server (``xql.serve``): each +test starts one in-process, registers its Datasets there, and queries +them through the ADBC Flight SQL driver. It runs once +``adbc-driver-flightsql`` is installed. Registration serves a Dataset in +place rather than ingesting it, so the tests of ingest modes and +temporary tables skip it. + ``XARRAY_SQL_TEST_ONLY`` (comma-separated names) restricts a run to those backends, e.g. one CI job per database. """ @@ -23,6 +30,9 @@ import uuid import pytest +import xarray as xr + +import xarray_sql as xql try: from adbc_driver_manager import dbapi @@ -79,6 +89,23 @@ class Backend: folds: bool = False """Whether unquoted names fold, so mixed-case ones need quotes.""" needs_uri: bool = False + served: bool = False + """Whether this is xarray-sql's own Flight SQL server, which serves a + registered Dataset in place instead of ingesting it.""" + + def open(self) -> "Database": + """Connect, starting a server first for the ``served`` backend.""" + if not self.served: + return Database(self, self.connect()) + if self.driver is None: + pytest.skip("adbc-driver-flightsql is not installed") + server = xql.serve({}) + try: + con = dataclasses.replace(self, uri=server.uri).connect() + except BaseException: + server.shutdown() + raise + return Database(self, con, server) def connect(self): if dbapi is None: @@ -177,17 +204,34 @@ def _credentials(name: str) -> tuple[tuple[str, str], ...]: options=_credentials("flightsql"), needs_uri=True, ), + Backend( + "served", + _module_driver("adbc_driver_flightsql"), + folds=True, + served=True, + ), ] class Database: """A connection plus the unique names a test creates, dropped after.""" - def __init__(self, backend: Backend, con) -> None: + def __init__( + self, + backend: Backend, + con, + server: xql.FlightSQLServer | None = None, + ) -> None: self.backend = backend self.con = con + self.server = server self._created: list[str] = [] + def register(self, name: str, ds: xr.Dataset, **kwargs): + """Register ``ds`` where this backend's queries will find it.""" + target = self.con if self.server is None else self.server + return xql.register(target, name, ds, **kwargs) + def name(self, base: str) -> str: """A fresh table (or schema) name, dropped when the test ends.""" name = f"{base}_{uuid.uuid4().hex[:8]}" @@ -203,7 +247,14 @@ def query(self, sql: str): cur.execute(sql) return cur + def close(self) -> None: + self.con.close() + if self.server is not None: + self.server.shutdown() + def cleanup(self) -> None: + if self.server is not None: + return # its tables go when the server shuts down postgresql = self.backend.name == "postgresql" if postgresql: self.con.rollback() diff --git a/tests/conftest.py b/tests/conftest.py index 323bd505..96f72918 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -4,7 +4,7 @@ import pandas as pd import xarray as xr -from ._adbc import BACKENDS, Database +from ._adbc import BACKENDS def rand_wx(start: str, end: str) -> xr.Dataset: @@ -156,7 +156,7 @@ def air_and_stations(): def db(request): """A connection to each available ADBC backend (see ``tests/_adbc.py``).""" backend = request.param - database = Database(backend, backend.connect()) + database = backend.open() yield database database.cleanup() - database.con.close() + database.close() diff --git a/tests/test_adbc_backend.py b/tests/test_adbc_backend.py index 378420fc..0b4d7cc2 100644 --- a/tests/test_adbc_backend.py +++ b/tests/test_adbc_backend.py @@ -4,7 +4,9 @@ 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. +``tests/_adbc.py``), including xarray-sql's own Flight SQL server, where +registering serves the Dataset in place; tests of one database's +specifics follow them. """ import numpy as np @@ -90,12 +92,17 @@ def _select_all(db, table: str): ) +def _ingests(db) -> None: + if db.backend.served: + pytest.skip("a served Dataset is registered in place, not ingested") + + # 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) + db.register(table, ds) out = xql.to_dataset(_select_all(db, table), template=ds) @@ -106,7 +113,7 @@ def test_round_trip_keeps_values_and_dtypes(db, ds): def test_aggregates_skip_missing_values(db, ds): table = db.name("weather") - xql.register(db.con, table, ds) + db.register(table, ds) cur = db.query( f"SELECT lat, lon, AVG(temperature) AS temperature FROM {table} " @@ -121,18 +128,20 @@ def test_aggregates_skip_missing_values(db, ds): def test_existing_table_is_not_overwritten_by_default(db, ds): + _ingests(db) table = db.name("weather") - xql.register(db.con, table, ds) + db.register(table, ds) with pytest.raises(dbapi.Error): - xql.register(db.con, table, ds) + db.register(table, ds) def test_replace_then_append(db, ds): + _ingests(db) 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") + db.register(table, ds) + db.register(table, ds.isel(time=slice(0, 4)), mode="replace") + db.register(table, ds.isel(time=slice(4, 8)), mode="append") out = xql.to_dataset(_select_all(db, table), template=ds) @@ -140,22 +149,24 @@ def test_replace_then_append(db, ds): def test_create_append_creates_then_appends(db, ds): + _ingests(db) table = db.name("weather") - xql.register(db.con, table, ds, mode="create_append") - xql.register(db.con, table, ds, mode="create_append") + db.register(table, ds, mode="create_append") + db.register(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): + _ingests(db) table = db.name("weather") if not db.backend.temporary: with pytest.raises(ValueError, match="temporary"): - xql.register(db.con, table, ds, temporary=True) + db.register(table, ds, temporary=True) return - xql.register(db.con, table, ds, temporary=True) + db.register(table, ds, temporary=True) temporary = f"{db.backend.temporary_prefix}{table}" count = db.query(f"SELECT COUNT(*) FROM {temporary}").fetchone()[0] @@ -165,11 +176,11 @@ def test_temporary_tables(db, ds): 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) + db.register(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) + db.register(name, mixed_ds, table_names=NAMES) table = f"{name}_atmosphere" cur = db.query( @@ -188,7 +199,7 @@ 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) + db.register(table, forecast) cur = db.query(f"SELECT step, lat, t2m FROM {table} ORDER BY step, lat") out = xql.to_dataset( @@ -201,7 +212,7 @@ def test_timedelta_coordinates_round_trip(db, forecast, chunks): def test_chunked_round_trip_spills_the_cursor(db, ds): table = db.name("weather") - xql.register(db.con, table, ds) + db.register(table, ds) cur = db.query( f"SELECT time, lat, lon, temperature FROM {table} " @@ -219,7 +230,7 @@ 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) + db.register(table, ds) start, end = ( db.backend.time_literal.format(f"2021-01-01 0{hour}:00:00") for hour in (4, 6) @@ -254,7 +265,7 @@ def test_subsecond_times_round_trip(db, chunks): {"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) + db.register(table, ds) cur = db.query(f"SELECT time, v FROM {table} ORDER BY time") out = xql.to_dataset( @@ -274,7 +285,7 @@ def test_awkward_variable_names_round_trip(db): coords={"x": [10, 20, 30]}, ).chunk({"x": 3}) table = db.name("awkward") - xql.register(db.con, table, ds) + db.register(table, ds) columns = ", ".join( db.quoted(n) for n in ["x", "select", "wind speed", "Order"] @@ -291,7 +302,7 @@ def test_text_coordinates_round_trip(db): coords={"station": ["O'Hare", 'say "hi"', "東京", "Zürich"]}, ).chunk({"station": 4}) table = db.name("stations") - xql.register(db.con, table, ds) + db.register(table, ds) cur = db.query(f"SELECT station, count FROM {table}") out = xql.to_dataset(cur, template=ds) @@ -313,7 +324,7 @@ def test_integer_extremes_round_trip(db): coords={"x": [0, 1, 2]}, ).chunk({"x": 3}) table = db.name("extremes") - xql.register(db.con, table, ds) + db.register(table, ds) cur = db.query(f"SELECT x, i64, u8, u32, u64 FROM {table} ORDER BY x") out = xql.to_dataset(cur, template=ds) @@ -328,7 +339,7 @@ def test_uint64_beyond_int64_is_never_silently_wrong(db): ).chunk({"x": 2}) table = db.name("huge") try: - xql.register(db.con, table, ds) + db.register(table, ds) except (ValueError, dbapi.Error): return # refused loudly: acceptable @@ -350,10 +361,10 @@ def test_nanosecond_times_are_kept_or_truncation_is_reported(db): table = db.name("nanos") if db.backend.microseconds: with pytest.warns(RuntimeWarning, match="microsecond"): - xql.register(db.con, table, ds) + db.register(table, ds) return - xql.register(db.con, table, ds) + db.register(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()) @@ -365,7 +376,7 @@ def test_many_chunks_ingest(db): coords={"time": np.arange(1000), "x": np.arange(200)}, ).chunk({"time": 100}) table = db.name("big") - xql.register(db.con, table, ds) + db.register(table, ds) count, total = db.query(f"SELECT COUNT(*), SUM(v) FROM {table}").fetchone() assert (count, float(total)) == (200_000, float(ds.v.sum())) @@ -373,7 +384,7 @@ def test_many_chunks_ingest(db): def test_empty_result_round_trips(db, ds): table = db.name("weather") - xql.register(db.con, table, ds) + db.register(table, ds) cur = db.query( f"SELECT time, lat, lon, temperature FROM {table} WHERE lat > 1000" @@ -386,7 +397,7 @@ def test_empty_result_round_trips(db, ds): 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) + db.register(table, ds) count = db.query(f"SELECT COUNT(*) FROM {table}").fetchone()[0] assert count == 8 * 5 * 6 @@ -396,9 +407,9 @@ 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) + db.register(table, ds) else: - xql.register(db.con, table, ds) + db.register(table, ds) count = db.query(f"SELECT COUNT(*) FROM {db.quoted(table)}").fetchone()[0] assert count == 8 * 5 * 6 @@ -416,8 +427,7 @@ 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.register( db.name("weather"), ds, ingest_options={"not.an.option": "x"}, @@ -433,7 +443,7 @@ def test_postgresql_schema_failure_explains_the_aborted_transaction( _only(db, "postgresql") with pytest.raises(RuntimeError, match="rollback"): - xql.register(db.con, db.name("pg_era5"), mixed_ds, table_names=NAMES) + db.register(db.name("pg_era5"), mixed_ds, table_names=NAMES) def test_postgresql_uses_an_existing_schema(db, mixed_ds): @@ -441,7 +451,7 @@ def test_postgresql_uses_an_existing_schema(db, mixed_ds): name = db.name("era5") db.query(f'CREATE SCHEMA "{name}"').close() - xql.register(db.con, name, mixed_ds, table_names=NAMES) + db.register(name, mixed_ds, table_names=NAMES) count = db.query(f"SELECT COUNT(*) FROM {name}.surface").fetchone()[0] assert count == 6 * 3 * 4 @@ -452,7 +462,7 @@ def test_postgresql_tables_have_statistics_after_register(db, ds): # registered table can pick plans that run for hours. _only(db, "postgresql") table = db.name("weather") - xql.register(db.con, table, ds) + db.register(table, ds) cur = db.query(f"SELECT reltuples FROM pg_class WHERE relname = '{table}'") assert cur.fetchone()[0] == 8 * 5 * 6 @@ -462,7 +472,7 @@ 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) + db.register(table, ds) def fail_on_second_chunk(block, block_info=None): if block_info[0]["chunk-location"][0] == 1: @@ -476,7 +486,7 @@ def fail_on_second_chunk(block, block_info=None): ) ) with pytest.raises((OSError, dbapi.Error)): - xql.register(db.con, table, broken, mode="replace") + db.register(table, broken, mode="replace") out = xql.to_dataset(_select_all(db, table), template=ds) xr.testing.assert_identical(out, ds.compute()) @@ -488,6 +498,6 @@ def test_mysql_keeps_the_default_database(db, mixed_ds): _only(db, "mysql", "mariadb") before = db.query("SELECT DATABASE()").fetchone()[0] - xql.register(db.con, db.name("era5"), mixed_ds, table_names=NAMES) + db.register(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 index 75f68668..435cde49 100644 --- a/tests/test_adbc_era5_integration.py +++ b/tests/test_adbc_era5_integration.py @@ -19,7 +19,7 @@ import xarray_sql as xql -from ._adbc import BACKENDS, Database +from ._adbc import BACKENDS pytestmark = pytest.mark.integration @@ -61,16 +61,16 @@ def era5() -> xr.Dataset: def db(request, era5): """Each backend, with the subset registered once for all its queries.""" backend = request.param - database = Database(backend, backend.connect()) + database = backend.open() name = database.name("era5") - xql.register(database.con, name, era5, table_names=NAMES) + database.register(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() + database.close() def test_area_mean_time_series(db, era5): diff --git a/tests/test_flight_sql.py b/tests/test_flight_sql.py new file mode 100644 index 00000000..183f8374 --- /dev/null +++ b/tests/test_flight_sql.py @@ -0,0 +1,342 @@ +"""Tests for serving Datasets over Arrow Flight SQL. + +A ``FlightSQLServer`` hosts lazily registered Datasets; the client here +is ADBC's Flight SQL driver, the same one remote users would connect +with, so the tests exercise the real wire protocol end to end. + +How served data round-trips (types, missing values, names, mixed +dimensions, the ARCO-ERA5 queries) is tested by the ADBC contract, where +the server is the ``served`` backend (``tests/_adbc.py``). The tests here +cover what only a server has: serving, discovery, read-only SQL, plain +Arrow Flight, and other databases as its clients. +""" + +import os +import threading +import urllib.request + +import numpy as np +import pandas as pd +import pyarrow as pa +import pyarrow.flight as flight +import pytest +import xarray as xr + +import xarray_sql as xql + +flight_sql = pytest.importorskip("adbc_driver_flightsql.dbapi") + + +@pytest.fixture +def ds() -> xr.Dataset: + np.random.seed(5) + return xr.Dataset( + data_vars=dict( + temperature=(["time", "lat", "lon"], np.random.randn(8, 5, 6)), + precipitation=(["time", "lat", "lon"], np.random.rand(8, 5, 6)), + ), + coords=dict( + time=pd.date_range("2021-01-01", periods=8, freq="h"), + lat=np.linspace(-10.0, 10.0, 5), + lon=np.linspace(0.0, 40.0, 6), + ), + attrs=dict(description="Synthetic weather."), + ).chunk({"time": 4}) + + +@pytest.fixture +def server(ds): + with xql.serve({"weather": ds}) as running: + yield running + + +@pytest.fixture +def con(server): + connection = flight_sql.connect(server.uri) + yield connection + connection.close() + + +def _query(con, sql: str): + cur = con.cursor() + cur.execute(sql) + return cur + + +def test_datasets_registered_while_serving_are_visible(server, con, ds): + server.register("later", ds[["precipitation"]]) + + count = _query(con, "SELECT COUNT(*) FROM later").fetchone()[0] + assert count == 8 * 5 * 6 + + +@pytest.mark.parametrize( + "statement", + [ + "CREATE EXTERNAL TABLE leak STORED AS CSV LOCATION '/etc/hosts'", + "COPY (SELECT 1) TO '/tmp/xarray-sql-flight-leak.csv'", + "CREATE TABLE copy AS SELECT * FROM weather", + "SET datafusion.execution.batch_size = 1", + # Read-only, but a file outside the served Datasets. + "SELECT * FROM '/etc/hosts'", + ], +) +def test_clients_cannot_reach_beyond_the_datasets(con, statement): + with pytest.raises(flight_sql.Error): + _query(con, statement).fetchall() + + +def test_a_malformed_descriptor_is_the_clients_error(server): + client = flight.FlightClient(server.uri) + with pytest.raises(pa.ArrowInvalid): + client.get_schema(flight.FlightDescriptor.for_command(b"not a command")) + + +def test_memory_limit_fails_the_query_not_the_server(): + ds = xr.Dataset( + {"v": (["time", "x"], np.random.rand(1_000, 200))}, + coords={"time": np.arange(1_000), "x": np.arange(200)}, + ).chunk({"time": 100}) + with xql.serve({"big": ds}, memory_limit=2**20) as server: + con = flight_sql.connect(server.uri) + with pytest.raises(flight_sql.Error, match="Resources exhausted"): + _query( + con, + "SELECT COUNT(*) FROM big a JOIN big b " + "ON a.time = b.time AND a.x = b.x", + ).fetchall() + count = _query(con, "SELECT COUNT(*) FROM big").fetchone()[0] + con.close() + + assert count == 1_000 * 200 + + +def test_a_failed_mixed_dimension_registration_adds_nothing(): + ds = xr.Dataset( + { + "t2m": (["time", "lat"], np.random.rand(4, 3)), + "temperature": (["time", "level", "lat"], np.random.rand(4, 2, 3)), + }, + coords={"time": np.arange(4), "lat": [0.0, 1.0, 2.0], "level": [1, 2]}, + ).chunk({"time": 2}) + surface, atmosphere = ("time", "lat"), ("time", "level", "lat") + server = xql.FlightSQLServer() + server.register( + "era5", ds, table_names={surface: "old", atmosphere: "atmosphere"} + ) + + # `surface` is new but `atmosphere` is taken: neither is added. + with pytest.raises(RuntimeError, match="atmosphere"): + server.register( + "era5", + ds, + table_names={surface: "surface", atmosphere: "atmosphere"}, + ) + with server.serve(): + con = flight_sql.connect(server.uri) + with pytest.raises(flight_sql.Error, match="surface"): + _query(con, "SELECT * FROM era5.surface").fetchall() + con.close() + + +def test_tables_are_discoverable(con): + objects = con.adbc_get_objects(depth="tables").read_all().to_pylist() + tables = { + table["table_name"] + for catalog in objects + for schema in catalog["catalog_db_schemas"] or [] + for table in schema["db_schema_tables"] or [] + } + + assert "weather" in tables + + +def test_shutdown_stops_the_server(ds): + with xql.serve({"weather": ds}) as server: + assert server.is_running + assert server.uri == f"grpc://127.0.0.1:{server.port}" + + assert not server.is_running + + +def test_shutdown_does_not_wait_forever_on_an_unread_result(): + big = xr.Dataset( + {"v": (["time", "x"], np.zeros((2_000, 1_000)))}, + coords={"time": np.arange(2_000), "x": np.arange(1_000)}, + ).chunk({"time": 100}) + server = xql.serve({"big": big}) + con = flight_sql.connect(server.uri) + cur = _query(con, "SELECT * FROM big") + cur.fetchone() # leave the rest of the stream unread + + done = threading.Event() + stopper = threading.Thread( + target=lambda: (server.shutdown(timeout=0.5), done.set()), + daemon=True, + ) + stopper.start() + stopper.join(timeout=30) + + assert done.is_set() + assert not server.is_running + + +def _read_path(server, *path): + client = flight.FlightClient(server.uri) + info = client.get_flight_info(flight.FlightDescriptor.for_path(*path)) + return client.do_get(info.endpoints[0].ticket).read_all() + + +def test_plain_flight_path_names_a_table(server, ds): + table = _read_path(server, "weather") + out = xql.to_dataset(table.sort_by([("time", "ascending")]), template=ds) + + xr.testing.assert_allclose(out, ds.compute()) + + +def test_plain_flight_path_finds_a_mixed_case_name(ds): + # ClickHouse's arrowFlight sends the name as given, unquoted. + server = xql.FlightSQLServer() + with pytest.warns(RuntimeWarning, match="quote"): + server.register("Weather", ds) + with server.serve(): + table = _read_path(server, "Weather") + + assert table.num_rows == 8 * 5 * 6 + + +def test_plain_flight_path_can_be_schema_qualified(ds): + upper = ds.temperature.expand_dims(level=[500, 850]).rename("upper") + server = xql.FlightSQLServer() + server.register( + "era5", + ds.assign(upper=upper), + table_names={("time", "lat", "lon"): "surface"}, + ) + with server.serve(): + table = _read_path(server, "era5.surface") + + assert table.num_rows == 8 * 5 * 6 + + +@pytest.mark.parametrize("via", ["server", "xql"]) +def test_mixed_case_warning_points_at_the_caller(ds, via): + server = xql.FlightSQLServer() + with pytest.warns(RuntimeWarning, match="quote") as record: + if via == "server": + server.register("Weather", ds) + else: + xql.register(server, "Weather", ds) + + assert record[0].filename == __file__ + + +def test_plain_flight_path_can_be_a_query(server, ds): + table = _read_path( + server, + "SELECT lat, lon, AVG(temperature) AS temperature FROM weather " + "GROUP BY lat, lon ORDER BY lat, lon", + ) + out = xql.to_dataset(table, template=ds) + + expected = ds.temperature.mean("time") + xr.testing.assert_allclose(out.temperature, expected.compute()) + + +def test_plain_flight_schema_of_a_path(server): + client = flight.FlightClient(server.uri) + result = client.get_schema(flight.FlightDescriptor.for_path("weather")) + + assert result.schema.names == [ + "time", + "lat", + "lon", + "temperature", + "precipitation", + ] + + +def test_clickhouse_reads_through_arrow_flight(ds): + uri = os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_URI", "") + if not uri.startswith(("http://", "https://")): + # chDB, which the other ClickHouse tests can use, is built without + # the arrowFlight table function. + pytest.skip( + "set XARRAY_SQL_TEST_CLICKHOUSE_URI to a ClickHouse server's " + "HTTP address to run against ClickHouse" + ) + # A ClickHouse in a container reaches this process through the Docker + # host's address, so the server listens on every interface then. + host = os.environ.get("XARRAY_SQL_TEST_CLICKHOUSE_FLIGHT_HOST") + bind = "0.0.0.0" if host else "127.0.0.1" + with xql.serve({"weather": ds}, host=bind) as server: + sql = ( + "SELECT round(avg(temperature), 9) FROM " + f"arrowFlight('{host or '127.0.0.1'}:{server.port}', 'weather')" + ) + request = urllib.request.Request(uri, data=sql.encode()) + with urllib.request.urlopen(request, timeout=60) as response: + avg = float(response.read().decode()) + + assert avg == pytest.approx(float(ds.temperature.mean())) + + +@pytest.fixture(scope="module") +def spark(): + jar = os.environ.get("XARRAY_SQL_TEST_FLIGHT_SQL_JDBC_JAR") + if not jar: + pytest.skip( + "set XARRAY_SQL_TEST_FLIGHT_SQL_JDBC_JAR to the Arrow Flight SQL " + "JDBC driver jar to run against Spark" + ) + pyspark_sql = pytest.importorskip("pyspark.sql") + # The JVM's zone applies to timestamps read over JDBC, so it must be + # UTC for xarray's (UTC) times to arrive unshifted. + session = ( + pyspark_sql.SparkSession.builder.master("local[1]") + .config("spark.jars", jar) + .config("spark.sql.session.timeZone", "UTC") + .config("spark.driver.extraJavaOptions", "-Duser.timezone=UTC") + .config("spark.ui.enabled", "false") + .getOrCreate() + ) + yield session + session.stop() + + +def _spark_read(spark, server, dbtable): + return ( + spark.read.format("jdbc") + .option( + "url", + f"jdbc:arrow-flight-sql://127.0.0.1:{server.port}" + "/?useEncryption=false", + ) + .option("driver", "org.apache.arrow.driver.jdbc.ArrowFlightJdbcDriver") + .option("dbtable", dbtable) + .load() + ) + + +def test_spark_reads_over_jdbc(spark, server, ds): + frame = _spark_read(spark, server, "weather").orderBy("time", "lat", "lon") + out = xql.to_dataset(frame.toArrow(), template=ds) + + xr.testing.assert_allclose(out, ds.compute()) + + +def test_spark_pushes_filters_to_the_server(spark, server, ds): + frame = ( + _spark_read(spark, server, "weather") + .where("time >= TIMESTAMP '2021-01-01 04:00:00' AND lat > 0") + .select("time", "lat", "lon", "temperature") + .orderBy("time", "lat", "lon") + ) + plan = frame._jdf.queryExecution().executedPlan().toString() + out = xql.to_dataset(frame.toArrow(), template=ds) + + assert "*GreaterThanOrEqual(time," in plan + expected = ds.temperature.isel(time=slice(4, None)).where( + ds.lat > 0, drop=True + ) + xr.testing.assert_allclose(out.temperature, expected.compute()) diff --git a/uv.lock b/uv.lock index 6e2321a4..614b82c0 100644 --- a/uv.lock +++ b/uv.lock @@ -9,6 +9,23 @@ resolution-markers = [ "python_full_version < '3.11'", ] +[[package]] +name = "adbc-driver-flightsql" +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/65/91/605b35b97aa5972a8026c0687e9996658d1d10ff7a0624173a0d0c7c5b93/adbc_driver_flightsql-1.12.0.tar.gz", hash = "sha256:300f67801ea016578e0c4bb798fc2b9024e405cfdd3f79fd77e37481a5767e45", size = 22097, upload-time = "2026-07-28T00:43:02.432Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7a/a4/5ad52abf0c4830585b86d68cff8ce07779cdee9b8cf390dff0ee21a29fa7/adbc_driver_flightsql-1.12.0-py3-none-macosx_10_15_x86_64.whl", hash = "sha256:c2c7209407c180d9d0dd94846203f184ab91e116782585bf29cb968c2686ba75", size = 8146345, upload-time = "2026-07-28T00:41:22.824Z" }, + { url = "https://files.pythonhosted.org/packages/61/36/49267a4fc8c07c066dd8c36b9c3d27265a4dc0fc478cb5acb7bd6a551d80/adbc_driver_flightsql-1.12.0-py3-none-macosx_11_0_arm64.whl", hash = "sha256:fb37efefceb7ceca64840c76c55e6280f3811b54dca305ec1781f9cc1a6da21e", size = 7552045, upload-time = "2026-07-28T00:41:27.147Z" }, + { url = "https://files.pythonhosted.org/packages/2b/5b/f8566aa6d05eb40225f5874f8f5f5e8b8b4a4ce866d817a5ff352f326302/adbc_driver_flightsql-1.12.0-py3-none-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:69742e4f8fcc254490aa7f8e254a3b704c24bde9d5d827a0926ea90a6b4c3a6d", size = 14981635, upload-time = "2026-07-28T00:41:33.2Z" }, + { url = "https://files.pythonhosted.org/packages/57/45/6468fb425f3c0ae868c1ca9ab8e15cba76b59e941edd7e80e892e5476728/adbc_driver_flightsql-1.12.0-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9adfb5e46abc77347476d09ee76dfc835fef2b4dcb0c76d08d6d35aa3a4cce17", size = 13699634, upload-time = "2026-07-28T00:41:38.232Z" }, + { url = "https://files.pythonhosted.org/packages/dc/c7/6773dedc54fad73277ee2d767dd817e203e949972b0cfb197df99e00b784/adbc_driver_flightsql-1.12.0-py3-none-win_amd64.whl", hash = "sha256:098f2498778fd3efd671fbeacba80987de98ebfeb0291da77fe0a26c6fa78532", size = 15033803, upload-time = "2026-07-28T00:41:42.483Z" }, +] + [[package]] name = "adbc-driver-manager" version = "1.12.0" @@ -2849,6 +2866,7 @@ polars = [ { name = "polars" }, ] test = [ + { name = "adbc-driver-flightsql" }, { name = "adbc-driver-manager" }, { name = "adbc-driver-sqlite" }, { name = "cftime" }, @@ -2872,6 +2890,7 @@ dev = [ [package.metadata] requires-dist = [ + { name = "adbc-driver-flightsql", marker = "extra == 'test'", specifier = ">=1.12" }, { 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'" }, diff --git a/xarray_sql/__init__.py b/xarray_sql/__init__.py index 1bbfe565..28672586 100644 --- a/xarray_sql/__init__.py +++ b/xarray_sql/__init__.py @@ -1,5 +1,11 @@ from . import cftime -from .backends import arrow_dataset, arrow_datasets, register +from .backends import ( + FlightSQLServer, + arrow_dataset, + arrow_datasets, + register, + serve, +) from .geometry import bbox_conjuncts from .df import from_map from .reader import read_xarray, read_xarray_table @@ -9,12 +15,14 @@ __all__ = [ "cftime", "XarrayContext", + "FlightSQLServer", "read_xarray_table", "read_xarray", "arrow_dataset", "arrow_datasets", "bbox_conjuncts", "register", + "serve", "to_dataset", "from_map", # deprecated ] diff --git a/xarray_sql/backends/__init__.py b/xarray_sql/backends/__init__.py index 476056e5..ef4926de 100644 --- a/xarray_sql/backends/__init__.py +++ b/xarray_sql/backends/__init__.py @@ -18,6 +18,7 @@ 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 .flight import FlightSQLServer, serve # also self-registers from .pyarrow import ( XarrayArrowStream, XarrayPushdownDataset, @@ -27,6 +28,7 @@ __all__ = [ "EngineAdapter", + "FlightSQLServer", "XarrayArrowStream", "XarrayPushdownDataset", "arrow_dataset", @@ -34,4 +36,5 @@ "get_adapter", "register", "register_adapter", + "serve", ] diff --git a/xarray_sql/backends/adbc.py b/xarray_sql/backends/adbc.py index 4a8c56bc..3ab3ec42 100644 --- a/xarray_sql/backends/adbc.py +++ b/xarray_sql/backends/adbc.py @@ -203,8 +203,13 @@ def _warn_on_lost_precision(ds: xr.Dataset, dialect: Dialect) -> None: ) -def _warn_on_folded_names(names: list[str], dialect: Dialect) -> None: - """Warn about names *dialect*'s database would fold if unquoted.""" +def _warn_on_folded_names( + names: list[str], dialect: Dialect, stacklevel: int = 4 +) -> None: + """Warn about names *dialect*'s database would fold if unquoted. + + *stacklevel* counts frames from here to the user's call. + """ if dialect.folds is None: return fold = str.lower if dialect.folds == "lower" else str.upper @@ -215,7 +220,7 @@ def _warn_on_folded_names(names: list[str], dialect: Dialect) -> None: f"{dialect.name} folds unquoted names to {dialect.folds}case, " f"so quote these in queries: {quoted}.", RuntimeWarning, - stacklevel=4, + stacklevel=stacklevel, ) diff --git a/xarray_sql/backends/base.py b/xarray_sql/backends/base.py index 0fd7ede3..b86dd8a0 100644 --- a/xarray_sql/backends/base.py +++ b/xarray_sql/backends/base.py @@ -65,8 +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, DuckDB, and ADBC DBAPI " - "connections." + f"Supported: DataFusion SessionContext, DuckDB, ADBC DBAPI " + "connections, and xarray_sql.FlightSQLServer." ) diff --git a/xarray_sql/backends/flight.py b/xarray_sql/backends/flight.py new file mode 100644 index 00000000..294cae6b --- /dev/null +++ b/xarray_sql/backends/flight.py @@ -0,0 +1,297 @@ +"""Serve lazy xarray Datasets over Arrow Flight SQL. + +The other adapters bring an engine to the data: the Dataset becomes a +table on a connection in this process. A +[FlightSQLServer][xarray_sql.backends.flight.FlightSQLServer] turns that +around and brings the data to remote clients. It hosts a DataFusion +session over the same lazy tables +[XarrayContext][xarray_sql.XarrayContext] uses and speaks +[Flight SQL](https://arrow.apache.org/docs/format/FlightSql.html), so any +Flight SQL client — ADBC's Flight SQL driver in Python, R, Go, or Java, +the Flight SQL JDBC and ODBC drivers, and the SQL tools built on them — +can query the Dataset without a copy:: + + server = xql.serve({"era5": ds}, port=8815) + + # elsewhere + import adbc_driver_flightsql.dbapi as flight_sql + con = flight_sql.connect("grpc://server-host:8815") + cur = con.cursor() + cur.execute("SELECT AVG(t2m) FROM era5 WHERE time >= '2020-01-01'") + +Queries keep partition pruning on dimension predicates and projection +pushdown: only the chunks and variables a query touches are read from +the source, only while it runs, and results stream back as Arrow record +batches. + +The server has no authentication or TLS, and binds to ``127.0.0.1`` +unless told otherwise. Its SQL is read-only: DDL, DML, and other +statements are rejected, so clients cannot read or write the server's +filesystem through it. +""" + +from __future__ import annotations + +import time +from typing import Any, TypeGuard + +import xarray as xr + +from .._native import FlightSqlServer as _NativeFlightSqlServer +from ..df import ( + Chunks, + TableNames, + group_vars_by_dims, + resolve_table_names, + shared_coord_arrays, +) +from ..reader import read_xarray_table +from ._adbc_dialects import DIALECTS +from .adbc import _warn_on_folded_names +from .base import register_adapter + +__all__ = ["FlightSQLServer", "serve"] + +_WILDCARD_HOSTS = ("", "0.0.0.0", "::") + + +class FlightSQLServer: + """A Flight SQL endpoint over lazily registered xarray Datasets. + + Register Datasets with + [register][xarray_sql.backends.flight.FlightSQLServer.register] (or + [xarray_sql.register][], which dispatches here), then start it with + [serve][xarray_sql.backends.flight.FlightSQLServer.serve]. Datasets + registered while the server runs are visible to the next query. + + The server runs on a background thread; the Python process must stay + alive while it serves. In a script, block with + [wait][xarray_sql.backends.flight.FlightSQLServer.wait]. It is also a + context manager that shuts down on exit. + """ + + def __init__(self, memory_limit: int | None = None) -> None: + """Create a server with no tables. + + Args: + memory_limit: Bytes that queries' sorts, joins, and + aggregations may hold at once. A query that needs more + spills to disk where it can and fails otherwise; the + server keeps serving. ``None`` (default) sets no limit. + Reading chunks isn't counted, so cap the process + externally too (e.g. a container memory limit) before + serving untrusted clients. + """ + self._native = _NativeFlightSqlServer(memory_limit) + self._host: str | None = None + self._port: int | None = None + + def register( + self, + name: str, + ds: xr.Dataset, + *, + chunks: Chunks = None, + table_names: TableNames = None, + ) -> FlightSQLServer: + """Register ``ds`` as a table named ``name``. + + A Dataset whose variables sit on different dimensions is split + into one table per dimension group, in a SQL schema named + ``name`` (``name.group``), exactly as + [XarrayContext.from_dataset][xarray_sql.XarrayContext.from_dataset] + registers it. + + Names are registered exactly as given. The server's SQL folds + unquoted names to lowercase, so a mixed-case name warns and must + be quoted in queries. + + Args: + name: The table name, or the schema name for a + mixed-dimension Dataset. + ds: An xarray Dataset. + chunks: Xarray-like chunks specification controlling partition + granularity. Defaults to the Dataset's existing chunks. + table_names: Maps a dimension group's exact dim tuple to the + name its table takes. + + Returns: + The server, to allow chaining. + """ + return self._register(name, ds, chunks, table_names, stacklevel=4) + + def _register( + self, + name: str, + ds: xr.Dataset, + chunks: Chunks, + table_names: TableNames, + stacklevel: int, + ) -> FlightSQLServer: + """``register``, warning *stacklevel* frames up (the user's call).""" + groups = group_vars_by_dims(ds) + names = resolve_table_names(ds, table_names) + _warn_on_folded_names( + [name] if len(groups) <= 1 else [name, *names.values()], + DIALECTS["datafusion"], + stacklevel=stacklevel, + ) + if len(groups) <= 1: + self._native.register_table(name, read_xarray_table(ds, chunks)) + return self + + # Every group's table is built before any is registered, so a + # failure leaves the server as it was. + coord_arrays = shared_coord_arrays(ds) + tables = [ + ( + names[dims], + read_xarray_table( + ds[var_names], chunks, coord_arrays=coord_arrays + ), + ) + for dims, var_names in groups.items() + ] + self._native.register_tables(name, tables) + return self + + def serve(self, host: str = "127.0.0.1", port: int = 0) -> FlightSQLServer: + """Start accepting connections on a background thread. + + Args: + host: The interface to bind. The default accepts only local + connections; ``"0.0.0.0"`` accepts them from anywhere the + network allows (the server has no authentication). + port: The TCP port. ``0`` (default) picks a free one; read it + back from [port][xarray_sql.backends.flight.FlightSQLServer.port]. + + Returns: + The server, to allow chaining. + """ + self._port = self._native.serve(host, port) + self._host = host + return self + + @property + def port(self) -> int: + """The TCP port the server is bound to.""" + if self._port is None: + raise RuntimeError("the server has not been started") + return self._port + + @property + def uri(self) -> str: + """A ``grpc://`` URI a local client can connect to.""" + host = self._host if self._host not in _WILDCARD_HOSTS else "localhost" + if host is not None and ":" in host: + host = f"[{host}]" + return f"grpc://{host}:{self.port}" + + @property + def is_running(self) -> bool: + """Whether the server is accepting connections.""" + return bool(self._native.is_running()) + + def wait(self, poll_interval: float = 0.5) -> None: + """Block until the server stops; Ctrl+C shuts it down. + + Raises: + RuntimeError: The server stopped on an error. + """ + try: + while self._native.is_running(): + time.sleep(poll_interval) + except KeyboardInterrupt: + pass + # Joins the server thread, raising what stopped it, if anything. + self.shutdown() + + def shutdown(self, timeout: float = 5.0) -> None: + """Stop accepting connections. + + Args: + timeout: Seconds in-flight queries may keep running, after + which their connections are closed. A client that stops + reading a result partway otherwise holds its stream open + indefinitely. + + Raises: + RuntimeError: The server had stopped on an error. + """ + self._native.shutdown(timeout) + + def __enter__(self) -> FlightSQLServer: + return self + + def __exit__(self, *exc: object) -> None: + self.shutdown() + + def __repr__(self) -> str: + state = f"serving on {self.uri}" if self.is_running else "stopped" + return f"FlightSQLServer({state})" + + +def serve( + datasets: dict[str, xr.Dataset], + host: str = "127.0.0.1", + port: int = 0, + *, + chunks: Chunks = None, + table_names: TableNames = None, + memory_limit: int | None = None, +) -> FlightSQLServer: + """Serve Datasets over Arrow Flight SQL, without copying them. + + Starts a [FlightSQLServer][xarray_sql.backends.flight.FlightSQLServer] + on a background thread with each Dataset registered under its key:: + + server = xql.serve({"era5": ds}, port=8815) + print(server.uri) # grpc://127.0.0.1:8815 + server.wait() # in a script: block until Ctrl+C + + Any Flight SQL client can then query the Datasets; with ADBC, + ``xql.to_dataset(cursor, template=ds)`` turns a result back into a + labeled Dataset on the client. + + Args: + datasets: Table name to Dataset. + host: The interface to bind; ``127.0.0.1`` by default. + port: The TCP port; ``0`` (default) picks a free one. + chunks: Chunks specification applied to every Dataset. + table_names: Dimension-group naming applied to every Dataset. + memory_limit: Bytes queries may hold at once; see + [FlightSQLServer][xarray_sql.backends.flight.FlightSQLServer]. + + Returns: + The running server. + """ + server = FlightSQLServer(memory_limit) + for name, ds in datasets.items(): + server.register(name, ds, chunks=chunks, table_names=table_names) + return server.serve(host, port) + + +@register_adapter +class FlightSQLAdapter: + """Registers Datasets on a [FlightSQLServer][xarray_sql.backends.flight.FlightSQLServer].""" + + @staticmethod + def matches(con: object) -> TypeGuard[FlightSQLServer]: + return isinstance(con, FlightSQLServer) + + @staticmethod + def register( + con: FlightSQLServer, + name: str, + ds: xr.Dataset, + *, + chunks: Chunks = None, + table_names: TableNames = None, + **kwargs: Any, + ) -> FlightSQLServer: + if kwargs: + raise TypeError( + f"unexpected options for a FlightSQLServer: {sorted(kwargs)}" + ) + # Called by xarray_sql.register, one frame further from the user. + return con._register(name, ds, chunks, table_names, stacklevel=5)