diff --git a/Cargo.lock b/Cargo.lock index 5cc827fc52..e00bfacf50 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -178,6 +178,8 @@ dependencies = [ "arrow-schema", "arrow-select", "flatbuffers", + "lz4_flex", + "zstd", ] [[package]] @@ -394,6 +396,18 @@ version = "1.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06" +[[package]] +name = "bigquery-benchmark-arrow-jobs-query" +version = "0.0.0" +dependencies = [ + "anyhow", + "clap", + "google-cloud-bigquery", + "google-cloud-wkt", + "stats_alloc", + "tokio", +] + [[package]] name = "bigquery-grpc-mock" version = "0.0.0" @@ -586,9 +600,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.4.4" +version = "1.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0ad534f4357a5264cce5019c989cf66a4f0dc4e0d1b1d15f8aacec0ff7360273" +checksum = "509591b7bcd67f4ef775afad7662703b4935daaa6ec0e5605cfb1090b32a2b6d" dependencies = [ "find-msvc-tools", "jobserver", @@ -861,9 +875,9 @@ dependencies = [ [[package]] name = "crypto-common" -version = "0.1.7" +version = "0.1.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +checksum = "1bfb12502f3fc46cca1bb51ac28df9d618d813cdc3d2f25b9fe775a34af26bb3" dependencies = [ "generic-array", "typenum", @@ -1024,7 +1038,7 @@ checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ "block-buffer 0.10.4", "const-oid 0.9.6", - "crypto-common 0.1.7", + "crypto-common 0.1.6", ] [[package]] @@ -1099,9 +1113,9 @@ dependencies = [ [[package]] name = "either" -version = "1.18.0" +version = "1.17.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "252afb9ae5eaa683babdc6a068b3f5726eb19e05070c731f9b2a23a7c3e8ed34" +checksum = "9e5e8f6c15a24b9a3ee5efec809ccd006d3b30e8b3bb63c39af737c7f87daa1d" [[package]] name = "elliptic-curve" @@ -1389,9 +1403,9 @@ dependencies = [ [[package]] name = "generic-array" -version = "0.14.7" +version = "0.14.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +checksum = "4bb6743198531e02858aeaea5398fcc883e71851fcbcb5a2f773e2fb6cb1edf2" dependencies = [ "typenum", "version_check", @@ -2327,6 +2341,7 @@ name = "google-cloud-bigquery" version = "0.16.1-preview" dependencies = [ "anyhow", + "arrow", "async-trait", "base64 0.23.1", "bigquery-grpc-mock", @@ -6614,9 +6629,9 @@ dependencies = [ [[package]] name = "h2" -version = "0.4.18" +version = "0.4.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "839c0e8a181239723652be9062bb56ca5bf5f64011f73b623f6f4fc59086a228" +checksum = "a9f37a958b41b3b19ee2707c06439c0e9e547e847223eb791ecb0cb821c65e27" dependencies = [ "atomic-waker", "bytes", @@ -6957,9 +6972,9 @@ checksum = "e590f038c1464a96894fd6d10127e90a8be4509f56ff7ecef851b15cee0b7caa" [[package]] name = "icu_provider" -version = "2.3.1" +version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d27bbb9d3abbefac45d55f647c9de1d44aafcd1186eb91879afef17c396c3e73" +checksum = "92a7ed671a6aad807a8651a2e1782a6598fda9ce5185dd8158549e95a91c6428" dependencies = [ "displaydoc", "icu_locale_core", @@ -7655,9 +7670,9 @@ dependencies = [ [[package]] name = "log" -version = "0.4.34" +version = "0.4.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f9f8bd3e56ce4dfc153cf470fffbfa98c7620958b312ca5c3a4b8d5181fd13c6" +checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" [[package]] name = "lru-slab" @@ -7665,6 +7680,15 @@ version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" +[[package]] +name = "lz4_flex" +version = "0.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ecbdfe44b1bd960b68170b417450a628c43f7cf56bb3c5317e61cb230ee7f226" +dependencies = [ + "twox-hash", +] + [[package]] name = "markdown" version = "1.0.0" @@ -8617,18 +8641,18 @@ dependencies = [ [[package]] name = "ref-cast" -version = "1.0.27" +version = "1.0.26" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7e440fb4e4b4147295338efb76001ab9e4efc0e5839df2c47fc5ac2381d365c3" +checksum = "216e8f773d7923bcba9ceb86a86c93cabb3903a11872fc3f138c49630e50b96d" dependencies = [ "ref-cast-impl", ] [[package]] name = "ref-cast-impl" -version = "1.0.27" +version = "1.0.26" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "92ecd8964f8453721699a1ed72037b0db49ce2f5a5138486ee89bed6f67cdf3a" +checksum = "2c9283685feec7d69af75fb0e858d5e7378f33fe4fc699383b2916ab9273e03c" dependencies = [ "proc-macro2", "quote", @@ -8939,9 +8963,9 @@ checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f" [[package]] name = "rustls-webpki" -version = "0.103.15" +version = "0.103.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f3c3cf1d8b1e7d4927e2d154c3fcb02979afb9939629c62cd9048d4f07b60ac2" +checksum = "0527518605e68109d875e248ea259b6758801cf165e4b2c2733ae3b51f12535a" dependencies = [ "aws-lc-rs", "ring", @@ -9434,6 +9458,12 @@ version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a2eb9349b6444b326872e140eb1cf5e7c522154d69e7a0ffb0fb81c06b37543f" +[[package]] +name = "stats_alloc" +version = "0.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c0e04424e733e69714ca1bbb9204c1a57f09f5493439520f9f68c132ad25eec" + [[package]] name = "storage-benchmark-appendable-object" version = "0.0.0" @@ -10199,6 +10229,12 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "twox-hash" +version = "2.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8464ec13c3691491391d9fce00f6416c9a48e46972f72d7865688be2080192c9" + [[package]] name = "typenum" version = "1.20.1" @@ -10299,9 +10335,9 @@ checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" [[package]] name = "uuid" -version = "1.25.0" +version = "1.24.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f053576934f05a761a402421fbbe3d425d9366f75f978806a037b3ca481abecc" +checksum = "2cefc03fd367c0c6d4305de1b312cf00248c4114f4a0418ce6a6af769e3b0bd9" dependencies = [ "getrandom 0.4.3", "js-sys", @@ -10751,9 +10787,9 @@ dependencies = [ [[package]] name = "zerovec" -version = "0.11.8" +version = "0.11.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bb0464e17806c1d976d5cba29399c7f08e516e279e2ba493f63123b5fca67dd8" +checksum = "94b5c6b5976d66c1d703c4fd17d3f5e43c8cedaacf604961b171adc7130896d8" dependencies = [ "yoke", "zerofrom", @@ -10762,9 +10798,9 @@ dependencies = [ [[package]] name = "zerovec-derive" -version = "0.11.6" +version = "0.11.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "34df6fc39dbd26ddc9c10e6a2984476e13acce22e64e4487636ef494369225da" +checksum = "9f212a141d820099d57ffafb9569be9617a6f27d3dc881fbee8fb56642f917a9" dependencies = [ "proc-macro2", "quote", @@ -10782,3 +10818,31 @@ name = "zmij" version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" + +[[package]] +name = "zstd" +version = "0.13.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e91ee311a569c327171651566e07972200e76fcfe2242a4fa446149a3881c08a" +dependencies = [ + "zstd-safe", +] + +[[package]] +name = "zstd-safe" +version = "7.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f49c4d5f0abb602a93fb8736af2a4f4dd9512e36f7f570d66e65ff867ed3b9d" +dependencies = [ + "zstd-sys", +] + +[[package]] +name = "zstd-sys" +version = "2.0.16+zstd.1.5.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91e19ebc2adc8f83e43039e79776e3fda8ca919132d68a1fed6a5faca2683748" +dependencies = [ + "cc", + "pkg-config", +] diff --git a/Cargo.toml b/Cargo.toml index cc7c36ec3e..11b22d3402 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -70,6 +70,7 @@ members = [ "src/bigquery", "src/bigquery-derive", "src/bigquery-read", + "src/bigquery/benchmarks/arrow_jobs_query", "src/bigquery/examples", "src/bigquery/grpc-mock", "src/bigtable", @@ -605,6 +606,7 @@ storage-grpc-mock = { path = "src/storage/grpc-mock" } [workspace.lints.rust] unexpected_cfgs = { level = "deny", check-cfg = [ + 'cfg(google_cloud_unstable_bigquery_arrow)', 'cfg(google_cloud_unstable_gapic_streaming)', 'cfg(google_cloud_unstable_grpc_rust)', 'cfg(google_cloud_unstable_grpc_server_streaming)', diff --git a/src/bigquery-derive/src/lib.rs b/src/bigquery-derive/src/lib.rs index fa0a90298c..1e230d51b5 100644 --- a/src/bigquery-derive/src/lib.rs +++ b/src/bigquery-derive/src/lib.rs @@ -112,6 +112,9 @@ pub fn derive_from_sql(input: TokenStream) -> TokenStream { let field_idents_struct_obj = fields .iter() .map(|f| f.ident.as_ref().expect("named field must have identifier")); + let field_idents_struct_arrow = fields + .iter() + .map(|f| f.ident.as_ref().expect("named field must have identifier")); let field_extractions_array = fields.iter().map(|f| { let field_name = f.ident.as_ref().expect("named field must have identifier"); @@ -133,6 +136,14 @@ pub fn derive_from_sql(input: TokenStream) -> TokenStream { } }); + let field_extractions_arrow = fields.iter().map(|f| { + let field_name = f.ident.as_ref().expect("named field must have identifier"); + let db_column_name = get_field_name(f); + quote! { + let #field_name = cell.take(#db_column_name)?; + } + }); + let expanded = quote! { impl google_cloud_bigquery::query::FromSql for #name { fn from_sql(value: wkt::Value) -> std::result::Result { @@ -156,6 +167,16 @@ pub fn derive_from_sql(input: TokenStream) -> TokenStream { }), } } + + fn from_arrow(cell: google_cloud_bigquery::query::ArrowCell<'_>) -> std::result::Result { + if cell.is_null() { + return std::result::Result::Err(google_cloud_bigquery::error::ConvertError::NotNull); + } + #( #field_extractions_arrow )* + std::result::Result::Ok(Self { + #( #field_idents_struct_arrow, )* + }) + } } }; diff --git a/src/bigquery/Cargo.toml b/src/bigquery/Cargo.toml index 3da282b651..4e1b04f79d 100644 --- a/src/bigquery/Cargo.toml +++ b/src/bigquery/Cargo.toml @@ -26,6 +26,7 @@ categories.workspace = true rust-version.workspace = true [dependencies] +arrow = { workspace = true, features = ["ipc", "ipc_compression"] } async-trait.workspace = true base64.workspace = true bytes.workspace = true diff --git a/src/bigquery/benchmarks/arrow_jobs_query/Cargo.toml b/src/bigquery/benchmarks/arrow_jobs_query/Cargo.toml new file mode 100644 index 0000000000..bb6838335d --- /dev/null +++ b/src/bigquery/benchmarks/arrow_jobs_query/Cargo.toml @@ -0,0 +1,33 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +[package] +name = "bigquery-benchmark-arrow-jobs-query" +version = "0.0.0" +publish = false +edition.workspace = true +authors.workspace = true +license.workspace = true +repository.workspace = true + +[lints] +workspace = true + +[dependencies] +anyhow.workspace = true +clap = { workspace = true, features = ["derive", "env", "help", "std"] } +google-cloud-bigquery = { workspace = true, features = ["default"] } +tokio = { workspace = true, features = ["full"] } +wkt.workspace = true +stats_alloc = "0.1" diff --git a/src/bigquery/benchmarks/arrow_jobs_query/src/main.rs b/src/bigquery/benchmarks/arrow_jobs_query/src/main.rs new file mode 100644 index 0000000000..230193b3c0 --- /dev/null +++ b/src/bigquery/benchmarks/arrow_jobs_query/src/main.rs @@ -0,0 +1,486 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use clap::{Parser, ValueEnum}; +use google_cloud_bigquery::client::BigQuery; +use google_cloud_bigquery::query::FromRow; +use stats_alloc::{Region, StatsAlloc}; +use std::alloc::System; +use std::time::{Duration, Instant}; + +#[global_allocator] +static GLOBAL: StatsAlloc = StatsAlloc::system(); + +fn format_bytes(bytes: usize) -> String { + if bytes >= 1024 * 1024 * 1024 { + format!("{:.2} GB", bytes as f64 / (1024.0 * 1024.0 * 1024.0)) + } else if bytes >= 1024 * 1024 { + format!("{:.2} MB", bytes as f64 / (1024.0 * 1024.0)) + } else if bytes >= 1024 { + format!("{:.2} KB", bytes as f64 / 1024.0) + } else { + format!("{bytes} B") + } +} + +// ============================================================================ +// CLI Arguments and Scenarios +// ============================================================================ + +#[derive(Parser, Debug)] +#[command( + name = "bigquery-benchmark-arrow-jobs-query", + about = "BigQuery jobs.query Benchmark: Arrow results format vs standard JSON (Speed & Allocations)" +)] +struct Args { + /// GCP Project ID (reads from GOOGLE_CLOUD_PROJECT if omitted). + #[arg(long, env = "GOOGLE_CLOUD_PROJECT")] + project_id: Option, + + /// Scenario to execute. + #[arg(long, value_enum, default_value = "synthetic-100k")] + scenario: Scenario, + + /// Custom query to run when scenario is `custom`. + #[arg(long)] + query: Option, + + /// Number of benchmark measurement iterations. + #[arg(long, default_value_t = 5)] + iterations: usize, + + /// Number of warmup iterations before measuring. + #[arg(long, default_value_t = 1)] + warmup: usize, + + /// Whether to enable BigQuery server-side query cache (default: false). + #[arg(long, default_value_t = false)] + use_query_cache: bool, + + /// Whether to deserialize rows into typed structs using FromRow. + #[arg(long, default_value_t = true)] + typed: bool, +} + +#[derive(ValueEnum, Clone, Copy, Debug, PartialEq, Eq)] +enum Scenario { + #[value(name = "synthetic-1")] + Synthetic1, + #[value(name = "synthetic-100")] + Synthetic100, + #[value(name = "synthetic-1k")] + Synthetic1k, + #[value(name = "synthetic-10k")] + Synthetic10k, + #[value(name = "synthetic-50k")] + Synthetic50k, + #[value(name = "synthetic-100k")] + Synthetic100k, + #[value(name = "synthetic-500k")] + Synthetic500k, + #[value(name = "wikipedia-10k")] + Wikipedia10k, + #[value(name = "wikipedia-100k")] + Wikipedia100k, + #[value(name = "custom")] + Custom, +} + +impl Scenario { + fn query(&self, custom_query: Option<&str>) -> String { + match self { + Scenario::Synthetic1 => Self::synthetic_query(1), + Scenario::Synthetic100 => Self::synthetic_query(100), + Scenario::Synthetic1k => Self::synthetic_query(1_000), + Scenario::Synthetic10k => Self::synthetic_query(10_000), + Scenario::Synthetic50k => Self::synthetic_query(50_000), + Scenario::Synthetic100k => Self::synthetic_query(100_000), + Scenario::Synthetic500k => Self::synthetic_query(500_000), + Scenario::Wikipedia10k => Self::wikipedia_query(10_000), + Scenario::Wikipedia100k => Self::wikipedia_query(100_000), + Scenario::Custom => custom_query + .expect("custom query must be provided when scenario is 'custom'") + .to_string(), + } + } + + fn synthetic_query(row_count: usize) -> String { + format!( + "SELECT \ + x AS id, \ + CONCAT('row_item_name_', CAST(x AS STRING)) AS name, \ + CAST(x AS FLOAT64) * 1.25 AS score, \ + (MOD(x, 2) = 0) AS is_even, \ + CURRENT_TIMESTAMP() AS created_at \ + FROM UNNEST(GENERATE_ARRAY(1, {row_count})) AS x" + ) + } + + fn wikipedia_query(limit: usize) -> String { + format!( + "SELECT \ + title, \ + id, \ + language, \ + wp_namespace, \ + is_redirect, \ + revision_id, \ + timestamp, \ + contributor_ip, \ + contributor_id, \ + contributor_username, \ + comment, \ + num_characters \ + FROM `bigquery-public-data.samples.wikipedia` \ + LIMIT {limit}" + ) + } + + fn is_wikipedia(&self) -> bool { + matches!(self, Scenario::Wikipedia10k | Scenario::Wikipedia100k) + } +} + +#[derive(FromRow, Debug, PartialEq)] +#[allow(dead_code)] +struct SyntheticRow { + id: i64, + name: String, + score: f64, + is_even: bool, + created_at: wkt::Timestamp, +} + +#[derive(FromRow, Debug, PartialEq)] +#[allow(dead_code)] +struct WikipediaRow { + title: Option, + id: Option, + language: Option, + wp_namespace: Option, + is_redirect: Option, + revision_id: Option, + timestamp: Option, + contributor_ip: Option, + contributor_id: Option, + contributor_username: Option, + comment: Option, + num_characters: Option, +} + +struct IterationResult { + query_duration: Duration, + read_duration: Duration, + total_duration: Duration, + rows_count: usize, + bytes_allocated: usize, + allocations_count: usize, +} + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + let args = Args::parse(); + + let project_id = args.project_id.ok_or_else(|| { + anyhow::anyhow!( + "Project ID must be provided via --project-id or GOOGLE_CLOUD_PROJECT env var" + ) + })?; + + let sql_query = args.scenario.query(args.query.as_deref()); + + println!( + "==========================================================================================================" + ); + println!(" BigQuery Query Benchmark (Speed & Memory)"); + println!( + "==========================================================================================================" + ); + #[cfg(google_cloud_unstable_bigquery_arrow)] + println!(" Arrow Acceleration: ENABLED (--cfg google_cloud_unstable_bigquery_arrow)"); + #[cfg(not(google_cloud_unstable_bigquery_arrow))] + println!(" Arrow Acceleration: DISABLED (Standard JSON mode)"); + println!(" Project ID: {project_id}"); + println!(" Scenario: {:?}", args.scenario); + println!(" Warmup Iterations: {}", args.warmup); + println!(" Measured Runs: {}", args.iterations); + println!(" Use Query Cache: {}", args.use_query_cache); + println!(" Typed Deserialization: {}", args.typed); + println!( + "==========================================================================================================" + ); + println!(); + + let client = BigQuery::builder() + .with_project_id(&project_id) + .build() + .await?; + + // Warmup runs + if args.warmup > 0 { + println!("Running {} warmup iteration(s)...", args.warmup); + for i in 1..=args.warmup { + print!(" Warmup {i}/{}: ", args.warmup); + let result = run_single_query( + &client, + &project_id, + &sql_query, + args.scenario, + args.use_query_cache, + args.typed, + ) + .await?; + println!( + "done ({} rows, query: {:.2?}, read: {:.2?}, total: {:.2?}, allocated: {})", + result.rows_count, + result.query_duration, + result.read_duration, + result.total_duration, + format_bytes(result.bytes_allocated) + ); + } + println!(); + } + + // Benchmark measured runs + println!("Running {} benchmark measurement(s)...", args.iterations); + println!( + "------------------------------------------------------------------------------------------------------------------" + ); + println!( + "{:<6} | {:<12} | {:<12} | {:<12} | {:<10} | {:<14} | {:<11} | {:<12}", + "Run", + "Query Time", + "Read/Iter", + "Total Time", + "Rows", + "Throughput", + "Allocated", + "Allocations" + ); + println!( + "------------------------------------------------------------------------------------------------------------------" + ); + + let mut results = Vec::with_capacity(args.iterations); + for i in 1..=args.iterations { + let result = run_single_query( + &client, + &project_id, + &sql_query, + args.scenario, + args.use_query_cache, + args.typed, + ) + .await?; + let rps = if result.read_duration.as_secs_f64() > 0.0 { + result.rows_count as f64 / result.read_duration.as_secs_f64() + } else { + 0.0 + }; + + println!( + "{:<6} | {:<12.2?} | {:<12.2?} | {:<12.2?} | {:<10} | {:>10.0} rows/s | {:<11} | {:>12}", + format!("#{i}"), + result.query_duration, + result.read_duration, + result.total_duration, + result.rows_count, + rps, + format_bytes(result.bytes_allocated), + result.allocations_count + ); + results.push(result); + } + println!( + "------------------------------------------------------------------------------------------------------------------" + ); + + // Print summary statistics + print_summary(&results); + + println!(); + println!("Tip: Compare Arrow vs JSON by running:"); + println!( + " Arrow: RUSTFLAGS=\"--cfg google_cloud_unstable_bigquery_arrow\" cargo run --release -p bigquery-benchmark-arrow-jobs-query -- --scenario {:?}", + args.scenario + ); + println!( + " JSON: cargo run --release -p bigquery-benchmark-arrow-jobs-query -- --scenario {:?}", + args.scenario + ); + println!(); + + Ok(()) +} + +async fn run_single_query( + client: &BigQuery, + project_id: &str, + query_str: &str, + scenario: Scenario, + use_query_cache: bool, + typed: bool, +) -> anyhow::Result { + let region = Region::new(&GLOBAL); + let start_total = Instant::now(); + + // 1. Submit query and wait until complete + let start_query = Instant::now(); + let complete_query = client + .query(query_str) + .with_project_id(project_id) + .set_use_query_cache(use_query_cache) + .until_done() + .await?; + let query_duration = start_query.elapsed(); + + // 2. Read and deserialize rows + let start_read = Instant::now(); + let mut iter = complete_query.read(); + let mut rows_count = 0; + + if typed { + if scenario.is_wikipedia() { + while let Some(row_res) = iter.next().await { + let row = row_res?; + let _typed_row: WikipediaRow = row.try_into()?; + rows_count += 1; + } + } else { + while let Some(row_res) = iter.next().await { + let row = row_res?; + let _typed_row: SyntheticRow = row.try_into()?; + rows_count += 1; + } + } + } else { + while let Some(row_res) = iter.next().await { + let _row = row_res?; + rows_count += 1; + } + } + let read_duration = start_read.elapsed(); + let total_duration = start_total.elapsed(); + + let stats = region.change(); + let bytes_allocated = stats.bytes_allocated; + let allocations_count = stats.allocations; + + Ok(IterationResult { + query_duration, + read_duration, + total_duration, + rows_count, + bytes_allocated, + allocations_count, + }) +} + +fn print_summary(results: &[IterationResult]) { + if results.is_empty() { + return; + } + + let n = results.len() as f64; + let query_times: Vec = results + .iter() + .map(|r| r.query_duration.as_secs_f64()) + .collect(); + let read_times: Vec = results + .iter() + .map(|r| r.read_duration.as_secs_f64()) + .collect(); + let total_times: Vec = results + .iter() + .map(|r| r.total_duration.as_secs_f64()) + .collect(); + let throughputs: Vec = results + .iter() + .map(|r| { + if r.read_duration.as_secs_f64() > 0.0 { + r.rows_count as f64 / r.read_duration.as_secs_f64() + } else { + 0.0 + } + }) + .collect(); + let bytes_allocs: Vec = results.iter().map(|r| r.bytes_allocated as f64).collect(); + let alloc_counts: Vec = results.iter().map(|r| r.allocations_count as f64).collect(); + + let avg = |v: &[f64]| v.iter().sum::() / n; + let min = |v: &[f64]| v.iter().cloned().fold(f64::INFINITY, f64::min); + let max = |v: &[f64]| v.iter().cloned().fold(f64::NEG_INFINITY, f64::max); + let std_dev = |v: &[f64], mean: f64| { + let variance = v.iter().map(|x| (x - mean).powi(2)).sum::() / n; + variance.sqrt() + }; + + let q_avg = avg(&query_times); + let r_avg = avg(&read_times); + let t_avg = avg(&total_times); + let tp_avg = avg(&throughputs); + let bytes_avg = avg(&bytes_allocs); + let count_avg = avg(&alloc_counts); + + let rows_avg = results[0].rows_count as f64; + let bytes_per_row = if rows_avg > 0.0 { + bytes_avg / rows_avg + } else { + 0.0 + }; + + println!("Summary Statistics (over {} runs):", results.len()); + println!( + " Query Execution Time: avg: {:.2?} (min: {:.2?}, max: {:.2?}, stddev: {:.2?})", + Duration::from_secs_f64(q_avg), + Duration::from_secs_f64(min(&query_times)), + Duration::from_secs_f64(max(&query_times)), + Duration::from_secs_f64(std_dev(&query_times, q_avg)) + ); + println!( + " Row Reading & Parsing: avg: {:.2?} (min: {:.2?}, max: {:.2?}, stddev: {:.2?})", + Duration::from_secs_f64(r_avg), + Duration::from_secs_f64(min(&read_times)), + Duration::from_secs_f64(max(&read_times)), + Duration::from_secs_f64(std_dev(&read_times, r_avg)) + ); + println!( + " Total End-to-End Time: avg: {:.2?} (min: {:.2?}, max: {:.2?}, stddev: {:.2?})", + Duration::from_secs_f64(t_avg), + Duration::from_secs_f64(min(&total_times)), + Duration::from_secs_f64(max(&total_times)), + Duration::from_secs_f64(std_dev(&total_times, t_avg)) + ); + println!( + " Row Throughput: avg: {:.0} rows/s (min: {:.0}, max: {:.0})", + tp_avg, + min(&throughputs), + max(&throughputs) + ); + println!( + " Total Heap Allocated: avg: {} ({:.1} bytes/row)", + format_bytes(bytes_avg as usize), + bytes_per_row + ); + println!( + " Total Allocations: avg: {:.0} allocs ({:.1} allocs/row)", + count_avg, + if rows_avg > 0.0 { + count_avg / rows_avg + } else { + 0.0 + } + ); +} diff --git a/src/bigquery/src/datatypes.rs b/src/bigquery/src/datatypes.rs index c47fa4595b..a4e23a5330 100644 --- a/src/bigquery/src/datatypes.rs +++ b/src/bigquery/src/datatypes.rs @@ -19,7 +19,7 @@ use crate::error::ConvertError; use crate::query::FromSql; -use crate::query::from_sql::parse_time; +use crate::query::from_sql::{ArrowCell, parse_time}; /// Represents a BigQuery time [INTERVAL] value. /// @@ -144,6 +144,35 @@ impl FromSql for Interval { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.downcast_value::(|arr, idx| { + let v = arr.value(idx); + let ym_sign = if v.months < 0 { -1 } else { 1 }; + let total_months = v.months.unsigned_abs(); + let years = (total_months / 12) as i32 * ym_sign; + let months = (total_months % 12) as i32 * ym_sign; + + let time_sign = if v.nanoseconds < 0 { -1 } else { 1 }; + let total_nanos = v.nanoseconds.unsigned_abs(); + let nanos = (total_nanos % 1_000_000_000) as i32 * time_sign; + let total_secs = total_nanos / 1_000_000_000; + let seconds = (total_secs % 60) as i32 * time_sign; + let total_mins = total_secs / 60; + let minutes = (total_mins % 60) as i32 * time_sign; + let hours = (total_mins / 60) as i32 * time_sign; + + Interval { + years, + months, + days: v.days, + hours, + minutes, + seconds, + nanos, + } + }) + } } /// Represents a BigQuery [RANGE] value. @@ -227,13 +256,33 @@ impl FromSql for Range { Ok(Range { start, end }) } + wkt::Value::Object(mut obj) => { + let start = match obj.remove("start") { + Some(wkt::Value::Null) | None => None, + Some(val) => Some(T::from_sql(val)?), + }; + let end = match obj.remove("end") { + Some(wkt::Value::Null) | None => None, + Some(val) => Some(T::from_sql(val)?), + }; + Ok(Range { start, end }) + } wkt::Value::Null => Err(ConvertError::NotNull), other => Err(ConvertError::TypeMismatch { - expected: "string", + expected: "string or object", got: other, }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + if cell.is_null() { + return Err(ConvertError::NotNull); + } + let start = cell.take("start")?; + let end = cell.take("end")?; + Ok(Range { start, end }) + } } #[cfg(test)] @@ -283,7 +332,7 @@ mod tests { #[test_case(wkt::Value::String("[UNBOUNDED, 2026-05-29)".to_string()) => Ok(Range { start: None, end: Some(google_cloud_type::model::Date::new().set_year(2026).set_month(5).set_day(29)) }) ; "date range unbounded start")] #[test_case(wkt::Value::String("[UNBOUNDED, UNBOUNDED)".to_string()) => Ok(Range { start: None, end: None }) ; "date range unbounded both")] #[test_case(wkt::Value::Null => Err(TestConvertError::NotNull) ; "null range")] - #[test_case(wkt::Value::Number(123.into()) => Err(TestConvertError::TypeMismatch("string")) ; "range type mismatch")] + #[test_case(wkt::Value::Number(123.into()) => Err(TestConvertError::TypeMismatch("string or object")) ; "range type mismatch")] #[test_case(wkt::Value::String("[2026-05-28)".to_string()) => Err(TestConvertError::Convert("invalid range format: expected 2 parts, got 1".to_string())) ; "range invalid format one part")] #[test_case(wkt::Value::String("[2026-05-28, 2026-05-29, 2026-05-30)".to_string()) => Err(TestConvertError::Convert("invalid range format: expected 2 parts, got 3".to_string())) ; "range invalid format three parts")] #[test_case(wkt::Value::String("[".to_string()) => Err(TestConvertError::Convert("invalid range format: missing enclosing brackets".to_string())) ; "range too short")] diff --git a/src/bigquery/src/query.rs b/src/bigquery/src/query.rs index 3427ec9c61..05f4f07e49 100644 --- a/src/bigquery/src/query.rs +++ b/src/bigquery/src/query.rs @@ -12,6 +12,7 @@ // See the License for the specific language governing permissions and // limitations under the License. +pub(super) mod arrow; pub(super) mod builder; pub(super) mod client; pub(super) mod client_builder; @@ -23,13 +24,14 @@ mod retry_policy; mod row; mod schema; +pub use arrow::ArrowCell; pub use iterator::RowIterator; pub use query_handle::{CompleteQuery, Query}; pub(crate) use schema::Schema; pub use from_sql::FromSql; pub use google_cloud_bigquery_derive::{FromRow, FromSql}; -pub use row::Row; +pub use row::{ColumnIndex, Row}; /// Result type for query execution. pub(super) type Result = std::result::Result; diff --git a/src/bigquery/src/query/arrow.rs b/src/bigquery/src/query/arrow.rs new file mode 100644 index 0000000000..72a02f8843 --- /dev/null +++ b/src/bigquery/src/query/arrow.rs @@ -0,0 +1,194 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use crate::error::ConvertError; +use crate::query::ColumnIndex; + +/// A reference to a single cell within an Arrow array. +#[doc(hidden)] +#[derive(Clone, Copy, Debug)] +pub struct ArrowCell<'a> { + array: &'a dyn arrow::array::Array, + pub(crate) row_idx: usize, +} + +impl<'a> ArrowCell<'a> { + /// Creates a new `ArrowCell`. + pub(crate) fn new(array: &'a dyn arrow::array::Array, row_idx: usize) -> Self { + Self { array, row_idx } + } + + /// Returns true if the cell is null. + #[doc(hidden)] + pub fn is_null(&self) -> bool { + self.array.is_null(self.row_idx) + } + + /// Extracts a field from a struct array cell by column index or name and converts it to `T`. + #[doc(hidden)] + pub fn take( + &self, + index: I, + ) -> Result { + let field_cell = self.struct_field_cell(&index)?; + T::from_arrow(field_cell) + } + + fn resolve_index( + &self, + col: &I, + struct_arr: &arrow::array::StructArray, + ) -> Result { + col.arrow_index(struct_arr) + .ok_or_else(|| ConvertError::MissingField(format!("{col}"))) + } + + /// Returns the data type of the underlying array. + pub(crate) fn data_type(&self) -> &arrow::datatypes::DataType { + self.array.data_type() + } + + /// Returns a string representation of the data type. + pub(crate) fn data_type_str(&self) -> String { + format!("{:?}", self.array.data_type()) + } + + /// Extracts a child `ArrowCell` from a struct array cell by column index or name. + fn struct_field_cell(&self, index: &I) -> Result, ConvertError> { + if self.is_null() { + return Err(ConvertError::NotNull); + } + + let Some(struct_arr) = self.downcast_ref::() else { + return Err(ConvertError::TypeMismatch { + expected: "struct array", + got: wkt::Value::String(self.data_type_str()), + }); + }; + + let idx = self.resolve_index(index, struct_arr)?; + let col = struct_arr.column(idx); + Ok(ArrowCell { + array: col.as_ref(), + row_idx: self.row_idx, + }) + } + + /// Downcasts the underlying Arrow array to a specific type. + pub(crate) fn downcast_ref(&self) -> Option<&T> { + self.array.as_any().downcast_ref::() + } + + /// Downcasts to the specified array type and returns the value at the cell's row index. + pub(crate) fn downcast_value(&self, f: F) -> Result + where + T: arrow::array::Array + 'static, + F: FnOnce(&T, usize) -> V, + { + if self.is_null() { + return Err(ConvertError::NotNull); + } + let arr = self + .downcast_ref::() + .ok_or_else(|| ConvertError::TypeMismatch { + expected: std::any::type_name::(), + got: wkt::Value::String(self.data_type_str()), + })?; + Ok(f(arr, self.row_idx)) + } + + /// Returns the cell's value as a boolean. + pub(crate) fn as_bool(&self) -> Result { + self.downcast_value::(|arr, idx| arr.value(idx)) + } + + /// Returns the cell's value as an `i64`. + pub(crate) fn as_i64(&self) -> Result { + self.downcast_value::(|arr, idx| arr.value(idx)) + } + + /// Returns the cell's value as an `i32`. + pub(crate) fn as_i32(&self) -> Result { + if self.is_null() { + return Err(ConvertError::NotNull); + } + if let Some(arr) = self.downcast_ref::() { + return i32::try_from(arr.value(self.row_idx)) + .map_err(|e| ConvertError::Convert(Box::new(e))); + } + if let Some(arr) = self.downcast_ref::() { + return Ok(arr.value(self.row_idx)); + } + Err(ConvertError::TypeMismatch { + expected: "Int64Array or Int32Array", + got: wkt::Value::String(self.data_type_str()), + }) + } + + /// Returns the cell's value as an `f64`. + pub(crate) fn as_f64(&self) -> Result { + self.downcast_value::(|arr, idx| arr.value(idx)) + } + + /// Returns the cell's value as an `f32`. + pub(crate) fn as_f32(&self) -> Result { + if self.is_null() { + return Err(ConvertError::NotNull); + } + if let Some(arr) = self.downcast_ref::() { + return Ok(arr.value(self.row_idx) as f32); + } + if let Some(arr) = self.downcast_ref::() { + return Ok(arr.value(self.row_idx)); + } + Err(ConvertError::TypeMismatch { + expected: "Float64Array or Float32Array", + got: wkt::Value::String(self.data_type_str()), + }) + } + + /// Returns the cell's value as a string slice (`&str`). + pub(crate) fn as_str(&self) -> Result<&str, ConvertError> { + if self.is_null() { + return Err(ConvertError::NotNull); + } + if let Some(arr) = self.downcast_ref::() { + Ok(arr.value(self.row_idx)) + } else if let Some(arr) = self.downcast_ref::() { + Ok(arr.value(self.row_idx)) + } else { + Err(ConvertError::TypeMismatch { + expected: "StringArray or LargeStringArray", + got: wkt::Value::String(self.data_type_str()), + }) + } + } + + /// Returns the cell's value as a byte slice (`&[u8]`). + pub(crate) fn as_bytes(&self) -> Result<&[u8], ConvertError> { + if self.is_null() { + return Err(ConvertError::NotNull); + } + if let Some(arr) = self.downcast_ref::() { + Ok(arr.value(self.row_idx)) + } else if let Some(arr) = self.downcast_ref::() { + Ok(arr.value(self.row_idx)) + } else { + Err(ConvertError::TypeMismatch { + expected: "BinaryArray or LargeBinaryArray", + got: wkt::Value::String(self.data_type_str()), + }) + } + } +} diff --git a/src/bigquery/src/query/execution.rs b/src/bigquery/src/query/execution.rs index f39bc44852..3d1735525b 100644 --- a/src/bigquery/src/query/execution.rs +++ b/src/bigquery/src/query/execution.rs @@ -20,7 +20,8 @@ use crate::query::retry_policy::JobRetryResult; use crate::query::{Query as QueryHandle, Result}; use google_cloud_bigquery_v2::client::JobService; use google_cloud_bigquery_v2::model::{ - InsertJobRequest, Job, JobConfiguration, PostQueryRequest, QueryRequest, QueryResponse, + DataFormatOptions, InsertJobRequest, Job, JobConfiguration, PostQueryRequest, QueryRequest, + QueryResponse, }; use google_cloud_gax::options::RequestOptionsBuilder as _; use google_cloud_gax::retry_state::RetryState; @@ -200,11 +201,20 @@ impl RetryContext { let query_request_id = generate_prefixed_id(QUERY_REQUEST_ID_PREFIX); let query_request: QueryRequest = self.template.request.clone().into(); let query_request = query_request - .set_format_options( - google_cloud_bigquery_v2::model::DataFormatOptions::new() - .set_use_int64_timestamp(true), - ) - .set_request_id(query_request_id); + .set_format_options(DataFormatOptions::new().set_use_int64_timestamp(true)); + #[cfg(google_cloud_unstable_bigquery_arrow)] + let query_request = query_request + .set_query_results_format(google_cloud_bigquery_v2::model::query_request::QueryResultsFormat::Arrow) + .set_results_format_serialization_options( + google_cloud_bigquery_v2::model::query_request::ResultsFormatSerializationOptions::ArrowSerializationOptions( + Box::new( + google_cloud_bigquery_v2::model::ArrowSerializationOptions::new().set_buffer_compression( + google_cloud_bigquery_v2::model::arrow_serialization_options::CompressionCodec::Zstd, + ), + ), + ), + ); + let query_request = query_request.set_request_id(query_request_id); let req = PostQueryRequest::new() .set_project_id(project_id) .set_query_request(query_request); @@ -263,7 +273,7 @@ mod tests { .clone() .expect("should have job_ref"); assert_eq!(job_ref.job_id, "my-job-123", "{job_ref:?}"); - assert!(query.cached_rows.is_some(), "{query:?}"); + assert!(query.cached_data.is_some(), "{query:?}"); Ok(()) } diff --git a/src/bigquery/src/query/from_sql.rs b/src/bigquery/src/query/from_sql.rs index f075bd30f7..4d4108ba8d 100644 --- a/src/bigquery/src/query/from_sql.rs +++ b/src/bigquery/src/query/from_sql.rs @@ -73,12 +73,111 @@ pub(crate) const BIGQUERY_DATETIME_SUBSEC_FORMAT: &[time::format_description::Fo pub trait FromSql: Sized { /// Converts a BigQuery `wkt::Value` into the implementing type. fn from_sql(value: wkt::Value) -> Result; + + /// Converts a BigQuery Arrow cell into the implementing type. + #[doc(hidden)] + fn from_arrow(cell: ArrowCell<'_>) -> Result { + let val = wkt::Value::from_arrow(cell)?; + Self::from_sql(val) + } } +pub use super::arrow::ArrowCell; + impl FromSql for wkt::Value { fn from_sql(value: wkt::Value) -> Result { Ok(value) } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + if cell.is_null() { + return Ok(wkt::Value::Null); + } + use arrow::datatypes::DataType; + match cell.data_type() { + DataType::Null => Ok(wkt::Value::Null), + DataType::Boolean => Ok(wkt::Value::Bool(cell.as_bool()?)), + DataType::Int64 => Ok(wkt::Value::Number(serde_json::Number::from(cell.as_i64()?))), + DataType::Float64 => { + let n = serde_json::Number::from_f64(cell.as_f64()?) + .ok_or_else(|| ConvertError::Convert("invalid f64 value".into()))?; + Ok(wkt::Value::Number(n)) + } + DataType::Utf8 | DataType::LargeUtf8 => { + Ok(wkt::Value::String(cell.as_str()?.to_string())) + } + DataType::Binary | DataType::LargeBinary => { + Ok(wkt::Value::String(BASE64_STANDARD.encode(cell.as_bytes()?))) + } + DataType::Date32 => { + let d = google_cloud_type::model::Date::from_arrow(cell)?; + Ok(wkt::Value::String(format!( + "{:04}-{:02}-{:02}", + d.year, d.month, d.day + ))) + } + DataType::Time64(arrow::datatypes::TimeUnit::Microsecond) => { + let t = google_cloud_type::model::TimeOfDay::from_arrow(cell)?; + if t.nanos == 0 { + Ok(wkt::Value::String(format!( + "{:02}:{:02}:{:02}", + t.hours, t.minutes, t.seconds + ))) + } else { + let subsec_micros = t.nanos / 1_000; + Ok(wkt::Value::String(format!( + "{:02}:{:02}:{:02}.{subsec_micros:06}", + t.hours, t.minutes, t.seconds + ))) + } + } + DataType::Timestamp(arrow::datatypes::TimeUnit::Microsecond, _) => { + let micros = cell.downcast_value::( + |arr, idx| arr.value(idx), + )?; + Ok(wkt::Value::String(micros.to_string())) + } + DataType::Decimal128(_, _) => { + let s = + cell.downcast_value::(|arr, idx| { + arr.value_as_string(idx) + })?; + Ok(wkt::Value::String(s)) + } + DataType::Decimal256(_, _) => { + let s = + cell.downcast_value::(|arr, idx| { + arr.value_as_string(idx) + })?; + Ok(wkt::Value::String(s)) + } + DataType::Interval(arrow::datatypes::IntervalUnit::MonthDayNano) => { + let interval = crate::datatypes::Interval::from_arrow(cell)?; + Ok(wkt::Value::String(format!( + "{}-{} {} {:02}:{:02}:{:02}.{:09}", + interval.years, + interval.months, + interval.days, + interval.hours, + interval.minutes, + interval.seconds, + interval.nanos + ))) + } + DataType::Struct(_) => { + let s = wkt::Struct::from_arrow(cell)?; + Ok(wkt::Value::Object(s)) + } + DataType::List(_) => { + let v = Vec::::from_arrow(cell)?; + Ok(wkt::Value::Array(v)) + } + _ => Err(ConvertError::TypeMismatch { + expected: "supported arrow value", + got: wkt::Value::String(cell.data_type_str()), + }), + } + } } impl FromSql for String { @@ -92,6 +191,10 @@ impl FromSql for String { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.as_str().map(ToString::to_string) + } } impl FromSql for i32 { @@ -111,6 +214,10 @@ impl FromSql for i32 { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.as_i32() + } } impl FromSql for i64 { @@ -129,6 +236,10 @@ impl FromSql for i64 { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.as_i64() + } } impl FromSql for f32 { @@ -148,6 +259,10 @@ impl FromSql for f32 { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.as_f32() + } } impl FromSql for f64 { @@ -166,6 +281,10 @@ impl FromSql for f64 { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.as_f64() + } } impl FromSql for bool { @@ -182,6 +301,10 @@ impl FromSql for bool { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.as_bool() + } } impl FromSql for Option { @@ -191,6 +314,14 @@ impl FromSql for Option { other => T::from_sql(other).map(Some), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + if cell.is_null() { + Ok(None) + } else { + T::from_arrow(cell).map(Some) + } + } } impl FromSql for Vec { @@ -204,6 +335,24 @@ impl FromSql for Vec { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + if cell.is_null() { + return Err(ConvertError::NotNull); + } + if let Some(arr) = cell.downcast_ref::() { + let value_arr = arr.value(cell.row_idx); + let mut result = Vec::with_capacity(value_arr.len()); + for i in 0..value_arr.len() { + result.push(T::from_arrow(ArrowCell::new(value_arr.as_ref(), i))?); + } + return Ok(result); + } + Err(ConvertError::TypeMismatch { + expected: "list array", + got: wkt::Value::String(cell.data_type_str()), + }) + } } impl FromSql for wkt::Struct { @@ -217,6 +366,25 @@ impl FromSql for wkt::Struct { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + if cell.is_null() { + return Err(ConvertError::NotNull); + } + if let Some(arr) = cell.downcast_ref::() { + let mut obj = wkt::Struct::new(); + let row_idx = cell.row_idx; + for (field, col) in arr.fields().iter().zip(arr.columns()) { + let val = wkt::Value::from_arrow(ArrowCell::new(col.as_ref(), row_idx))?; + obj.insert(field.name().clone(), val); + } + return Ok(obj); + } + Err(ConvertError::TypeMismatch { + expected: "struct array", + got: wkt::Value::String(cell.data_type_str()), + }) + } } impl FromSql for wkt::Timestamp { @@ -241,6 +409,14 @@ impl FromSql for wkt::Timestamp { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + let micros = + cell.downcast_value::(|arr, idx| { + arr.value(idx) + })?; + timestamp_from_micros(micros) + } } fn timestamp_from_micros(micros: i64) -> Result { @@ -269,6 +445,18 @@ impl FromSql for google_cloud_type::model::Date { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + let days = + cell.downcast_value::(|arr, idx| arr.value(idx))?; + let date = time::OffsetDateTime::from_unix_timestamp(days as i64 * 86400) + .map_err(|e| ConvertError::Convert(Box::new(e)))? + .date(); + Ok(google_cloud_type::model::Date::new() + .set_year(date.year()) + .set_month(u8::from(date.month()) as i32) + .set_day(date.day() as i32)) + } } pub(crate) fn parse_time(s: &str) -> Result { @@ -298,6 +486,24 @@ impl FromSql for google_cloud_type::model::TimeOfDay { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + let micros = + cell.downcast_value::(|arr, idx| { + arr.value(idx) + })?; + let nanos = (micros % 1_000_000) * 1_000; + let total_secs = micros / 1_000_000; + let seconds = total_secs % 60; + let total_mins = total_secs / 60; + let minutes = total_mins % 60; + let hours = total_mins / 60; + Ok(google_cloud_type::model::TimeOfDay::new() + .set_hours(hours as i32) + .set_minutes(minutes as i32) + .set_seconds(seconds as i32) + .set_nanos(nanos as i32)) + } } impl FromSql for google_cloud_type::model::DateTime { @@ -327,6 +533,23 @@ impl FromSql for google_cloud_type::model::DateTime { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + let micros = + cell.downcast_value::(|arr, idx| { + arr.value(idx) + })?; + let odt = time::OffsetDateTime::from_unix_timestamp_nanos(micros as i128 * 1_000) + .map_err(|e| ConvertError::Convert(Box::new(e)))?; + Ok(google_cloud_type::model::DateTime::new() + .set_year(odt.year()) + .set_month(u8::from(odt.month()) as i32) + .set_day(odt.day() as i32) + .set_hours(odt.hour() as i32) + .set_minutes(odt.minute() as i32) + .set_seconds(odt.second() as i32) + .set_nanos(odt.nanosecond() as i32)) + } } impl FromSql for google_cloud_type::model::Decimal { @@ -343,6 +566,25 @@ impl FromSql for google_cloud_type::model::Decimal { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + if cell.is_null() { + return Err(ConvertError::NotNull); + } + let row_idx = cell.row_idx; + if let Some(arr) = cell.downcast_ref::() { + let s = arr.value_as_string(row_idx); + return Ok(google_cloud_type::model::Decimal::new().set_value(s)); + } + if let Some(arr) = cell.downcast_ref::() { + let s = arr.value_as_string(row_idx); + return Ok(google_cloud_type::model::Decimal::new().set_value(s)); + } + Err(ConvertError::TypeMismatch { + expected: "decimal", + got: wkt::Value::String(cell.data_type_str()), + }) + } } impl FromSql for rust_decimal::Decimal { @@ -371,6 +613,39 @@ impl FromSql for rust_decimal::Decimal { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + if cell.is_null() { + return Err(ConvertError::NotNull); + } + let row_idx = cell.row_idx; + if let Some(arr) = cell.downcast_ref::() { + let val = arr.value(row_idx); + let scale = arr.scale() as u32; + return rust_decimal::Decimal::try_from_i128_with_scale(val, scale) + .map_err(|e| ConvertError::Convert(Box::new(e))); + } + if let Some(arr) = cell.downcast_ref::() { + let s = arr.value_as_string(row_idx); + let trimmed = if let Some((int_part, frac_part)) = s.split_once('.') { + let frac_trimmed = frac_part.trim_end_matches('0'); + if frac_trimmed.is_empty() { + int_part.to_string() + } else { + format!("{int_part}.{frac_trimmed}") + } + } else { + s + }; + return trimmed + .parse::() + .map_err(|e| ConvertError::Convert(Box::new(e))); + } + Err(ConvertError::TypeMismatch { + expected: "decimal", + got: wkt::Value::String(cell.data_type_str()), + }) + } } impl FromSql for Vec { @@ -386,12 +661,20 @@ impl FromSql for Vec { }), } } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.as_bytes().map(|b| b.to_vec()) + } } impl FromSql for bytes::Bytes { fn from_sql(value: wkt::Value) -> Result { Vec::::from_sql(value).map(bytes::Bytes::from) } + + fn from_arrow(cell: ArrowCell<'_>) -> Result { + cell.as_bytes().map(bytes::Bytes::copy_from_slice) + } } #[cfg(test)] diff --git a/src/bigquery/src/query/iterator.rs b/src/bigquery/src/query/iterator.rs index 9495e020e8..3d9bbe7c9b 100644 --- a/src/bigquery/src/query/iterator.rs +++ b/src/bigquery/src/query/iterator.rs @@ -13,10 +13,14 @@ // limitations under the License. use crate::error::RowError; +use crate::query::query_handle::CachedData; use crate::query::{CompleteQuery, Row, Schema}; +use arrow::ipc::reader::StreamReader; +use arrow::record_batch::RecordBatch; use google_cloud_bigquery_v2::client::JobService; use google_cloud_bigquery_v2::model::{GetQueryResultsRequest, JobReference}; use std::collections::VecDeque; +use std::io::{Cursor, Read}; use std::sync::Arc; pub type Result = std::result::Result; @@ -53,18 +57,41 @@ pub struct RowIterator { job_ref: Option, schema: Arc, page_token: Option, + record_batches: VecDeque>, + row_index: usize, rows: VecDeque, max_results: Option, } impl RowIterator { pub(crate) fn new(q: CompleteQuery) -> Self { + let (rows, record_batches) = match q.cached_data { + CachedData::Rows(rows) => (rows, VecDeque::new()), + CachedData::Arrow { + serialized_record_batch, + serialized_schema, + } => { + let reader = StreamReader::try_new( + Cursor::new(serialized_schema).chain(Cursor::new(serialized_record_batch)), + None, + ) + .expect("valid arrow IPC stream"); // TODO: convert error + let batches = reader + .map(|res| res.map(Arc::new)) + .collect::, _>>() + .expect("valid record batches"); // TODO: convert error + (VecDeque::new(), batches) + } + }; + Self { job_service: q.job_service, job_ref: q.job_ref, schema: q.schema, page_token: q.page_token, - rows: q.cached_rows, + record_batches, + row_index: 0, + rows, max_results: q.max_results, } } @@ -111,6 +138,16 @@ impl RowIterator { /// ``` pub async fn next(&mut self) -> Option> { loop { + while let Some(batch) = self.record_batches.front() { + if self.row_index < batch.num_rows() { + let idx = self.row_index; + self.row_index += 1; + return Some(Row::try_new_from_arrow(batch, idx, &self.schema)); + } + self.record_batches.pop_front(); + self.row_index = 0; + } + if let Some(raw_row) = self.rows.pop_front() { return Some(Row::try_new(raw_row, &self.schema)); } @@ -119,7 +156,7 @@ impl RowIterator { return Some(Err(e)); } - if self.rows.is_empty() && self.page_token.is_none() { + if self.record_batches.is_empty() && self.rows.is_empty() && self.page_token.is_none() { return None; } } @@ -417,4 +454,60 @@ mod tests { ); Ok(()) } + + #[tokio::test] + async fn test_row_iterator_cached_arrow() -> TestResult { + use arrow::array::{Int64Array, StringArray}; + use arrow::datatypes::{DataType, Field, Schema as ArrowSchema}; + use arrow::ipc::writer::StreamWriter; + + let arrow_schema = Arc::new(ArrowSchema::new(vec![ + Field::new("col", DataType::Utf8, false), + Field::new("num", DataType::Int64, false), + ])); + + let mut schema_buf = Vec::new(); + let _ = StreamWriter::try_new(&mut schema_buf, &arrow_schema)?; + + let col = StringArray::from(vec!["hello", "world"]); + let num = Int64Array::from(vec![42, 100]); + let batch = RecordBatch::try_new(arrow_schema.clone(), vec![Arc::new(col), Arc::new(num)])?; + + let mut batch_buf = Vec::new(); + let mut writer = StreamWriter::try_new(&mut batch_buf, &arrow_schema)?; + writer.write(&batch)?; + let batch_buf = batch_buf[schema_buf.len()..].to_vec(); + + let table_schema = TableSchema::new().set_fields([ + TableFieldSchema::new().set_name("col").set_type("STRING"), + TableFieldSchema::new().set_name("num").set_type("INTEGER"), + ]); + let schema = Arc::new(Schema::new(table_schema)); + + let job_service = create_job_service(MockJobService::new()); + let q = CompleteQuery { + job_service, + job_ref: None, + cached_data: CachedData::Arrow { + serialized_schema: schema_buf.into(), + serialized_record_batch: batch_buf.into(), + }, + schema, + page_token: None, + metadata: crate::generated::CompleteQueryMetadata::default(), + max_results: None, + }; + + let mut iter = q.read(); + let row1 = iter.next().await.expect("row 1")?; + assert_eq!(row1.get::("col"), "hello"); + assert_eq!(row1.get::("num"), 42); + + let row2 = iter.next().await.expect("row 2")?; + assert_eq!(row2.get::("col"), "world"); + assert_eq!(row2.get::("num"), 100); + + assert!(iter.next().await.is_none()); + Ok(()) + } } diff --git a/src/bigquery/src/query/query_handle.rs b/src/bigquery/src/query/query_handle.rs index edbe559359..74ab1e0f07 100644 --- a/src/bigquery/src/query/query_handle.rs +++ b/src/bigquery/src/query/query_handle.rs @@ -17,8 +17,10 @@ use crate::generated::{CompleteQueryMetadata, QueryMetadata}; use crate::query::execution::RetryContext; use crate::query::retry_policy::JobRetryResult; use crate::query::{Result, RowIterator, Schema}; +use bytes::Bytes; use google_cloud_bigquery_v2::builder::job_service::GetJob; use google_cloud_bigquery_v2::client::JobService; +use google_cloud_bigquery_v2::model::query_response::{Results, ResultsSchema}; use google_cloud_bigquery_v2::model::{ GetQueryResultsRequest, GetQueryResultsResponse, Job, JobReference, QueryResponse, }; @@ -56,11 +58,20 @@ pub struct Query { pub(crate) job_service: Arc, pub(crate) completed: bool, pub(crate) metadata: QueryMetadata, - pub(crate) cached_rows: Option>, + pub(crate) cached_data: Option, pub(crate) max_results: Option, pub(crate) retry_context: Option, } +#[derive(Clone, Debug)] +pub(crate) enum CachedData { + Rows(VecDeque), + Arrow { + serialized_record_batch: Bytes, + serialized_schema: Bytes, + }, +} + impl Query { pub(crate) fn from_job( job_service: Arc, @@ -76,7 +87,7 @@ impl Query { Self { job_service, completed, - cached_rows: None, + cached_data: None, metadata: QueryMetadata::from(initial_job), retry_context, max_results, @@ -90,12 +101,26 @@ impl Query { max_results: Option, ) -> Self { let completed = query_response.job_complete.unwrap_or(false); - let cached_rows = VecDeque::from(std::mem::take(&mut query_response.rows)); + let cached_data = if let ( + Some(ResultsSchema::ArrowSchema(schema)), + Some(Results::ArrowRecordBatch(results)), + ) = ( + query_response.results_schema.take(), + query_response.results.take(), + ) { + Some(CachedData::Arrow { + serialized_record_batch: results.serialized_record_batch, + serialized_schema: schema.serialized_schema, + }) + } else { + let cached_rows = VecDeque::from(std::mem::take(&mut query_response.rows)); + Some(CachedData::Rows(cached_rows)) + }; let metadata = QueryMetadata::from(query_response); Self { job_service, completed, - cached_rows: Some(cached_rows), + cached_data, metadata, retry_context, max_results, @@ -201,16 +226,16 @@ impl Query { job_service, completed, metadata, - cached_rows, + cached_data, max_results, retry_context, } = self; - if completed && let Some(cached_rows) = cached_rows { + if let (true, Some(cached_data)) = (completed, cached_data) { return Ok(CompleteQuery::from_query_metadata( job_service, metadata, - cached_rows, + cached_data, max_results, )); } @@ -280,7 +305,7 @@ impl Query { pub struct CompleteQuery { pub(crate) job_service: Arc, pub(crate) job_ref: Option, - pub(crate) cached_rows: VecDeque, + pub(crate) cached_data: CachedData, pub(crate) schema: Arc, pub(crate) page_token: Option, pub(crate) metadata: CompleteQueryMetadata, @@ -307,7 +332,7 @@ impl CompleteQuery { Self { job_service, job_ref: Some(job_ref.clone()), - cached_rows, + cached_data: CachedData::Rows(cached_rows), page_token, schema, metadata, @@ -318,14 +343,31 @@ impl CompleteQuery { pub(crate) fn from_query_metadata( job_service: Arc, metadata: QueryMetadata, - cached_rows: VecDeque, + cached_data: CachedData, max_results: Option, ) -> Self { let job_ref = metadata.job_reference.clone(); let metadata = CompleteQueryMetadata::from(metadata); - // DDL/DML queries have no schema. - let schema = metadata.schema.clone().unwrap_or_default(); - let schema = Arc::new(Schema::new(schema)); + let schema = match &cached_data { + CachedData::Rows(_) => { + // DDL/DML queries have no schema. + let schema = metadata.schema.clone().unwrap_or_default(); + Arc::new(Schema::new(schema)) + } + CachedData::Arrow { + serialized_schema, .. + } => { + match Schema::try_from_arrow_ipc(serialized_schema) { + Ok(s) => Arc::new(s), + Err(_) => { + // DDL/DML queries have no schema. + let schema = metadata.schema.clone().unwrap_or_default(); + Arc::new(Schema::new(schema)) + } + } + } + }; + let page_token = if metadata.page_token.is_empty() { None } else { @@ -334,7 +376,7 @@ impl CompleteQuery { Self { job_service, job_ref, - cached_rows, + cached_data, page_token, schema, metadata, @@ -517,7 +559,7 @@ mod tests { mut query_res: QueryResponse, max_results: Option, ) -> Self { - let cached_rows = std::mem::take(&mut query_res.rows).into(); + let cached_rows = CachedData::Rows(VecDeque::from(std::mem::take(&mut query_res.rows))); let metadata = QueryMetadata::from(query_res); Self::from_query_metadata(job_service, metadata, cached_rows, max_results) } @@ -542,7 +584,10 @@ mod tests { let completed = query.until_done().await?; assert_eq!(completed.job_ref.as_ref().unwrap().job_id, "some_job_id"); assert_eq!(completed.page_token, Some("some_page_token".to_string())); - assert_eq!(completed.cached_rows.len(), 1); + match &completed.cached_data { + CachedData::Rows(rows) => assert_eq!(rows.len(), 1), + _ => panic!("expected rows"), + } let metadata = completed.metadata(); assert_eq!(metadata.cache_hit, Some(true)); @@ -604,7 +649,10 @@ mod tests { let completed = query.until_done().await?; assert_eq!(completed.job_ref.as_ref().unwrap().job_id, "some_job_id"); assert_eq!(completed.page_token, None); - assert_eq!(completed.cached_rows.len(), 2); + match &completed.cached_data { + CachedData::Rows(rows) => assert_eq!(rows.len(), 2), + _ => panic!("expected rows"), + } let metadata = completed.metadata(); assert_eq!(metadata.cache_hit, Some(false)); diff --git a/src/bigquery/src/query/row.rs b/src/bigquery/src/query/row.rs index 84709f8a1f..13dd096ce8 100644 --- a/src/bigquery/src/query/row.rs +++ b/src/bigquery/src/query/row.rs @@ -13,7 +13,9 @@ // limitations under the License. use crate::error::{ConvertError, RowError}; +use crate::query::from_sql::ArrowCell; use crate::query::{FromSql, Schema}; +use arrow::record_batch::RecordBatch; use google_cloud_bigquery_v2::model::TableFieldSchema; use std::sync::Arc; use wkt::{ListValue, Struct, Value}; @@ -62,10 +64,19 @@ pub type Result = std::result::Result; /// ``` #[derive(Clone, Debug)] pub struct Row { - pub(crate) values: Value, + pub(crate) inner: RowInner, pub(crate) schema: Arc, } +#[derive(Clone, Debug)] +pub(crate) enum RowInner { + Json(ListValue), + Arrow { + batch: Arc, + row_idx: usize, + }, +} + mod sealed { /// A sealed trait to prevent external implementation of `ColumnIndex`. pub trait ColumnIndex {} @@ -80,24 +91,43 @@ mod sealed { pub trait ColumnIndex: sealed::ColumnIndex + std::fmt::Display { /// Returns the index of the column in the given row, if it exists. fn index(&self, row: &Row) -> Option; + + /// Returns the index of the column in the given arrow struct array, if it exists. + fn arrow_index(&self, struct_arr: &arrow::array::StructArray) -> Option; } impl ColumnIndex for usize { fn index(&self, row: &Row) -> Option { row.schema.get_field_by_index(*self).map(|_| *self) } + + fn arrow_index(&self, struct_arr: &arrow::array::StructArray) -> Option { + if *self < struct_arr.num_columns() { + Some(*self) + } else { + None + } + } } impl ColumnIndex for &str { fn index(&self, row: &Row) -> Option { row.schema.get_field_index_by_name(self) } + + fn arrow_index(&self, struct_arr: &arrow::array::StructArray) -> Option { + struct_arr.fields().iter().position(|f| f.name() == self) + } } impl ColumnIndex for String { fn index(&self, row: &Row) -> Option { self.as_str().index(row) } + + fn arrow_index(&self, struct_arr: &arrow::array::StructArray) -> Option { + self.as_str().arrow_index(struct_arr) + } } impl Row { @@ -105,7 +135,29 @@ impl Row { let values = convert_row(row, schema.fields())?; Ok(Self { - values: Value::Array(values), + inner: RowInner::Json(values), + schema: schema.clone(), + }) + } + + pub(crate) fn try_new_from_arrow( + batch: &Arc, + row_idx: usize, + schema: &Arc, + ) -> Result { + if batch.num_columns() != schema.len() { + return Err(RowError::InvalidRowFormat(format!( + "schema and row cell mismatch (expected {}, got {})", + schema.len(), + batch.num_columns() + ))); + } + + Ok(Self { + inner: RowInner::Arrow { + batch: Arc::clone(batch), + row_idx, + }, schema: schema.clone(), }) } @@ -155,15 +207,35 @@ impl Row { /// ``` pub fn try_get(&self, index: I) -> Result { let idx = self.resolve_index(&index)?; - let val = self - .values - .get(idx) - .ok_or_else(|| RowError::IndexOutOfRange { - index: idx, - len: self.schema.len(), - })?; - - self.convert_value_at(idx, val.clone()) + match &self.inner { + RowInner::Json(values) => { + let val = values.get(idx).ok_or_else(|| RowError::IndexOutOfRange { + index: idx, + len: self.schema.len(), + })?; + self.convert_value_at(idx, val.clone()) + } + RowInner::Arrow { batch, row_idx } => { + let col = batch + .columns() + .get(idx) + .ok_or_else(|| RowError::IndexOutOfRange { + index: idx, + len: self.schema.len(), + })?; + T::from_arrow(ArrowCell::new(col.as_ref(), *row_idx)).map_err(|e| { + let field_name = self + .schema + .get_field_by_index(idx) + .map(|f| f.name.clone()) + .unwrap_or_else(|| idx.to_string()); + RowError::TypeConversion { + column: field_name, + source: e, + } + }) + } + } } /// Takes ownership of a value from the row by column name or zero-based @@ -189,18 +261,40 @@ impl Row { /// ``` pub fn take(&mut self, index: I) -> Result { let idx = self.resolve_index(&index)?; - - let val = self - .values - .get_mut(idx) - .ok_or_else(|| RowError::IndexOutOfRange { - index: idx, - len: self.schema.len(), - })?; - - // swap out the value in-place to avoid clones - let owned_val = std::mem::replace(val, Value::Null); - self.convert_value_at(idx, owned_val) + match &mut self.inner { + RowInner::Json(values) => { + let val = values + .get_mut(idx) + .ok_or_else(|| RowError::IndexOutOfRange { + index: idx, + len: self.schema.len(), + })?; + + // swap out the value in-place to avoid clones + let owned_val = std::mem::replace(val, Value::Null); + self.convert_value_at(idx, owned_val) + } + RowInner::Arrow { batch, row_idx } => { + let col = batch + .columns() + .get(idx) + .ok_or_else(|| RowError::IndexOutOfRange { + index: idx, + len: self.schema.len(), + })?; + T::from_arrow(ArrowCell::new(col.as_ref(), *row_idx)).map_err(|e| { + let field_name = self + .schema + .get_field_by_index(idx) + .map(|f| f.name.clone()) + .unwrap_or_else(|| idx.to_string()); + RowError::TypeConversion { + column: field_name, + source: e, + } + }) + } + } } /// Retrieves a value from the row by column name or zero-based index. @@ -832,4 +926,169 @@ mod tests { assert!(matches!(err, RowError::ColumnNotFound(col) if col == "custom_int")); Ok(()) } + + #[test] + fn try_new_from_arrow_batch() -> TestResult { + use arrow::array::{ + BooleanArray, Float64Array, Int64Array, StringArray, TimestampMicrosecondArray, + }; + use arrow::datatypes::{DataType, Field, Schema as ArrowSchema, TimeUnit}; + + let arrow_schema = Arc::new(ArrowSchema::new(vec![ + Field::new("name", DataType::Utf8, false), + Field::new("age", DataType::Int64, true), + Field::new("active", DataType::Boolean, false), + Field::new("score", DataType::Float64, false), + Field::new( + "created_ts", + DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())), + false, + ), + Field::new( + "created_dt", + DataType::Timestamp(TimeUnit::Microsecond, None), + false, + ), + ])); + + let name = StringArray::from(vec!["Alice", "Bob"]); + let age = Int64Array::from(vec![Some(30), None]); + let active = BooleanArray::from(vec![true, false]); + let score = Float64Array::from(vec![98.5, 87.25]); + let created_ts = + TimestampMicrosecondArray::from(vec![1_600_000_000_000_000, 1_700_000_000_000_000]) + .with_timezone("UTC"); + let created_dt = + TimestampMicrosecondArray::from(vec![1_600_000_000_000_000, 1_700_000_000_000_000]); + + let batch = Arc::new(RecordBatch::try_new( + arrow_schema, + vec![ + Arc::new(name), + Arc::new(age), + Arc::new(active), + Arc::new(score), + Arc::new(created_ts), + Arc::new(created_dt), + ], + )?); + + let table_schema = TableSchema::new().set_fields([ + TableFieldSchema::new().set_name("name").set_type("STRING"), + TableFieldSchema::new().set_name("age").set_type("INTEGER"), + TableFieldSchema::new() + .set_name("active") + .set_type("BOOLEAN"), + TableFieldSchema::new().set_name("score").set_type("FLOAT"), + TableFieldSchema::new() + .set_name("created_ts") + .set_type("TIMESTAMP"), + TableFieldSchema::new() + .set_name("created_dt") + .set_type("DATETIME"), + ]); + let schema = Arc::new(Schema::new(table_schema)); + + let row0 = Row::try_new_from_arrow(&batch, 0, &schema)?; + assert_eq!(row0.get::("name"), "Alice"); + assert_eq!(row0.get::, _>("age"), Some(30)); + assert!(row0.get::("active")); + assert_eq!(row0.get::("score"), 98.5); + assert_eq!( + row0.get::("created_ts"), + wkt::Timestamp::new(1_600_000_000, 0).unwrap() + ); + + let row1 = Row::try_new_from_arrow(&batch, 1, &schema)?; + assert_eq!(row1.get::("name"), "Bob"); + assert_eq!(row1.get::, _>("age"), None); + assert!(!row1.get::("active")); + assert_eq!(row1.get::("score"), 87.25); + assert_eq!( + row1.get::("created_ts"), + wkt::Timestamp::new(1_700_000_000, 0).unwrap() + ); + + Ok(()) + } + + #[test] + fn try_new_from_arrow_interval() -> TestResult { + use crate::datatypes::Interval; + use arrow::array::IntervalMonthDayNanoArray; + use arrow::datatypes::{DataType, Field, IntervalUnit, Schema as ArrowSchema}; + + let arrow_schema = Arc::new(ArrowSchema::new(vec![Field::new( + "duration", + DataType::Interval(IntervalUnit::MonthDayNano), + false, + )])); + + let intervals = IntervalMonthDayNanoArray::from(vec![ + arrow::datatypes::IntervalMonthDayNanoType::make_value( + 14, + 3, + (4 * 3600 + 5 * 60 + 6) * 1_000_000_000 + 789_123_456, + ), + arrow::datatypes::IntervalMonthDayNanoType::make_value( + -14, + -3, + -((4 * 3600 + 5 * 60 + 6) * 1_000_000_000 + 123_000_000), + ), + arrow::datatypes::IntervalMonthDayNanoType::make_value(i32::MIN, 0, i64::MIN), + ]); + + let batch = Arc::new(RecordBatch::try_new( + arrow_schema, + vec![Arc::new(intervals)], + )?); + + let table_schema = TableSchema::new().set_fields([TableFieldSchema::new() + .set_name("duration") + .set_type("INTERVAL")]); + let schema = Arc::new(Schema::new(table_schema)); + + let row0 = Row::try_new_from_arrow(&batch, 0, &schema)?; + let int0: Interval = row0.get("duration"); + assert_eq!( + int0, + Interval { + years: 1, + months: 2, + days: 3, + hours: 4, + minutes: 5, + seconds: 6, + nanos: 789_123_456, + } + ); + + let row1 = Row::try_new_from_arrow(&batch, 1, &schema)?; + let int1: Interval = row1.get("duration"); + assert_eq!( + int1, + Interval { + years: -1, + months: -2, + days: -3, + hours: -4, + minutes: -5, + seconds: -6, + nanos: -123_000_000, + } + ); + + // Verifies no overflow on i32::MIN and i64::MIN + let row2 = Row::try_new_from_arrow(&batch, 2, &schema)?; + let int2: Interval = row2.get("duration"); + assert_eq!(int2.years, -178956970); + assert_eq!(int2.months, -8); + assert_eq!(int2.days, 0); + assert_eq!(int2.hours, -2562047); + assert_eq!(int2.minutes, -47); + assert_eq!(int2.seconds, -16); + assert_eq!(int2.nanos, -854_775_808); + + Ok(()) + } } diff --git a/src/bigquery/src/query/schema.rs b/src/bigquery/src/query/schema.rs index 26f3173c85..30ff473103 100644 --- a/src/bigquery/src/query/schema.rs +++ b/src/bigquery/src/query/schema.rs @@ -12,7 +12,10 @@ // See the License for the specific language governing permissions and // limitations under the License. +use arrow::datatypes::DataType; +use arrow::ipc::reader::StreamReader; use google_cloud_bigquery_v2::model::{TableFieldSchema, TableSchema}; +use std::io::Cursor; /// Schema of a table. #[derive(Clone, Debug)] @@ -38,4 +41,110 @@ impl Schema { pub(crate) fn fields(&self) -> &[TableFieldSchema] { &self.0.fields } + + pub(crate) fn try_from_arrow_ipc( + serialized_schema: &[u8], + ) -> Result { + let reader = StreamReader::try_new(Cursor::new(serialized_schema), None).map_err(|e| { + crate::error::RowError::InvalidRowFormat(format!("failed to parse arrow schema: {e}")) + })?; + let table_schema = table_schema_from_arrow_schema(&reader.schema()); + Ok(Self(table_schema)) + } +} + +fn table_schema_from_arrow_schema(arrow_schema: &arrow::datatypes::Schema) -> TableSchema { + let fields: Vec = arrow_schema + .fields() + .iter() + .map(|f| arrow_field_to_table_field(f)) + .collect(); + TableSchema::new().set_fields(fields) +} + +fn arrow_field_to_table_field(field: &arrow::datatypes::Field) -> TableFieldSchema { + let tf = TableFieldSchema::new().set_name(field.name().clone()); + let mode = if field.is_nullable() { + "NULLABLE" + } else { + "REQUIRED" + }; + + match field.data_type() { + DataType::Boolean => tf.set_type("BOOLEAN").set_mode(mode), + DataType::Int8 + | DataType::Int16 + | DataType::Int32 + | DataType::Int64 + | DataType::UInt8 + | DataType::UInt16 + | DataType::UInt32 + | DataType::UInt64 => tf.set_type("INTEGER").set_mode(mode), + DataType::Float32 | DataType::Float64 => tf.set_type("FLOAT64").set_mode(mode), + DataType::Utf8 | DataType::LargeUtf8 => tf.set_type("STRING").set_mode(mode), + DataType::Binary | DataType::LargeBinary => tf.set_type("BYTES").set_mode(mode), + DataType::Date32 | DataType::Date64 => tf.set_type("DATE").set_mode(mode), + DataType::Time32(_) | DataType::Time64(_) => tf.set_type("TIME").set_mode(mode), + DataType::Timestamp(_, _) => tf.set_type("TIMESTAMP").set_mode(mode), + DataType::Interval(_) => tf.set_type("INTERVAL").set_mode(mode), + DataType::Decimal128(_, _) | DataType::Decimal256(_, _) => { + tf.set_type("NUMERIC").set_mode(mode) + } + DataType::Struct(fields) => { + let sub_fields: Vec = fields + .iter() + .map(|f| arrow_field_to_table_field(f)) + .collect(); + tf.set_type("RECORD").set_mode(mode).set_fields(sub_fields) + } + DataType::List(sub_field) | DataType::LargeList(sub_field) => { + let mut sub = arrow_field_to_table_field(sub_field); + sub.name = field.name().clone(); + sub.mode = "REPEATED".to_string(); + sub + } + _ => tf.set_type("STRING").set_mode(mode), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::datatypes::{DataType, Field, Schema as ArrowSchema, TimeUnit}; + + #[test] + fn test_from_arrow_schema() { + let arrow_schema = ArrowSchema::new(vec![ + Field::new("name", DataType::Utf8, false), + Field::new("age", DataType::Int64, true), + Field::new( + "tags", + DataType::List(std::sync::Arc::new(Field::new( + "item", + DataType::Utf8, + true, + ))), + true, + ), + Field::new( + "created", + DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())), + false, + ), + ]); + + let table_schema = table_schema_from_arrow_schema(&arrow_schema); + let schema = Schema::new(table_schema); + assert_eq!(schema.len(), 4); + assert_eq!(schema.get_field_index_by_name("name"), Some(0)); + assert_eq!(schema.get_field_index_by_name("age"), Some(1)); + assert_eq!(schema.get_field_index_by_name("tags"), Some(2)); + assert_eq!(schema.get_field_index_by_name("created"), Some(3)); + + let f0 = schema.get_field_by_index(0).unwrap(); + assert_eq!(f0.name, "name"); + + let f2 = schema.get_field_by_index(2).unwrap(); + assert_eq!(f2.name, "tags"); + } } diff --git a/tests/bigquery/Cargo.toml b/tests/bigquery/Cargo.toml index c7f852da78..2dcd9903b7 100644 --- a/tests/bigquery/Cargo.toml +++ b/tests/bigquery/Cargo.toml @@ -31,7 +31,7 @@ log-integration-tests = [] [dependencies] anyhow.workspace = true -arrow = { workspace = true, features = ["ipc"] } +arrow = { workspace = true, features = ["ipc", "ipc_compression"] } bigquery-samples = { workspace = true } bytes.workspace = true futures.workspace = true