paimon-datafusion 0.3.0

Apache Paimon DataFusion Integration
Documentation
// Licensed to the Apache Software Foundation (ASF) under one
// or more contributor license agreements.  See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership.  The ASF licenses this file
// to you 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
//
//   http://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.

//! Shared test helpers for DataFusion integration tests.

use std::sync::Arc;

use datafusion::arrow::array::{Int32Array, LargeStringArray, StringArray, StringViewArray};
use paimon::{CatalogOptions, FileSystemCatalog, Options};
use paimon_datafusion::SQLContext;
use tempfile::TempDir;

use arrow_array::{Array, RecordBatch, UInt64Array};

#[allow(dead_code)]
pub fn string_value(array: &dyn Array, row: usize) -> &str {
    if let Some(array) = array.as_any().downcast_ref::<StringArray>() {
        array.value(row)
    } else if let Some(array) = array.as_any().downcast_ref::<LargeStringArray>() {
        array.value(row)
    } else if let Some(array) = array.as_any().downcast_ref::<StringViewArray>() {
        array.value(row)
    } else {
        panic!("expected a string array, got {}", array.data_type())
    }
}

pub fn create_test_env() -> (TempDir, Arc<FileSystemCatalog>) {
    let temp_dir = TempDir::new().expect("Failed to create temp dir");
    let warehouse = format!("file://{}", temp_dir.path().display());
    let mut options = Options::new();
    options.set(CatalogOptions::WAREHOUSE, warehouse);
    let catalog = FileSystemCatalog::new(options).expect("Failed to create catalog");
    (temp_dir, Arc::new(catalog))
}

pub async fn create_sql_context(catalog: Arc<FileSystemCatalog>) -> SQLContext {
    let mut ctx = SQLContext::new();
    ctx.register_catalog("paimon", catalog).await.unwrap();
    ctx
}

#[allow(dead_code)]
pub async fn setup_sql_context() -> (TempDir, SQLContext) {
    let (tmp, catalog) = create_test_env();
    let sql_context = create_sql_context(catalog).await;
    sql_context
        .sql("CREATE SCHEMA paimon.test_db")
        .await
        .expect("CREATE SCHEMA failed");
    (tmp, sql_context)
}

#[allow(dead_code)]
pub async fn collect_id_name(sql_context: &SQLContext, sql: &str) -> Vec<(i32, String)> {
    let mut rows = collect_id_name_in_batch_order(sql_context, sql).await;
    rows.sort_by_key(|(id, _)| *id);
    rows
}

#[allow(dead_code)]
pub async fn collect_id_name_in_batch_order(
    sql_context: &SQLContext,
    sql: &str,
) -> Vec<(i32, String)> {
    let batches = sql_context.sql(sql).await.unwrap().collect().await.unwrap();
    collect_id_name_from_batches_in_order(&batches)
}

#[allow(dead_code)]
pub fn collect_id_name_from_batches_in_order(batches: &[RecordBatch]) -> Vec<(i32, String)> {
    let mut rows = Vec::new();
    for batch in batches {
        let ids = batch
            .column_by_name("id")
            .and_then(|c| c.as_any().downcast_ref::<Int32Array>())
            .expect("id column");
        let names = batch.column_by_name("name").expect("name column");
        for i in 0..batch.num_rows() {
            rows.push((ids.value(i), string_value(names.as_ref(), i).to_string()));
        }
    }
    rows
}

#[allow(dead_code)]
pub async fn collect_id_value(sql_context: &SQLContext, sql: &str) -> Vec<(i32, i32)> {
    let batches = sql_context.sql(sql).await.unwrap().collect().await.unwrap();
    let mut rows = Vec::new();
    for batch in &batches {
        let ids = batch
            .column_by_name("id")
            .and_then(|c| c.as_any().downcast_ref::<Int32Array>())
            .expect("id column");
        let vals = batch
            .column_by_name("value")
            .and_then(|c| c.as_any().downcast_ref::<Int32Array>())
            .expect("value column");
        for i in 0..batch.num_rows() {
            rows.push((ids.value(i), vals.value(i)));
        }
    }
    rows.sort_by_key(|(id, _)| *id);
    rows
}

#[allow(dead_code)]
pub async fn row_count(sql_context: &SQLContext, sql: &str) -> usize {
    let batches = sql_context.sql(sql).await.unwrap().collect().await.unwrap();
    batches.iter().map(|b| b.num_rows()).sum()
}

/// Execute SQL and collect results, discarding the output.
#[allow(dead_code)]
pub async fn exec(sql_context: &SQLContext, s: &str) {
    sql_context.sql(s).await.unwrap().collect().await.unwrap();
}

/// Extract the count from a DML result (returns a single UInt64 column).
#[allow(dead_code)]
pub async fn dml_count(sql_context: &SQLContext, sql_str: &str) -> u64 {
    let result = sql_context
        .sql(sql_str)
        .await
        .unwrap()
        .collect()
        .await
        .unwrap();
    result[0]
        .column(0)
        .as_any()
        .downcast_ref::<UInt64Array>()
        .unwrap()
        .value(0)
}

#[allow(dead_code)]
pub async fn assert_sql_error(sql_context: &SQLContext, sql: &str, expected_substring: &str) {
    let err_msg = match sql_context.sql(sql).await {
        Ok(df) => match df.collect().await {
            Ok(_) => panic!("Expected error containing '{expected_substring}', but got Ok"),
            Err(err) => err.to_string(),
        },
        Err(err) => err.to_string(),
    };
    assert!(
        err_msg.contains(expected_substring),
        "Error message '{err_msg}' does not contain '{expected_substring}'"
    );
}

/// Collect (i32, i32, String) rows from batches, sorted by (col0, col1).
#[allow(dead_code)]
pub fn collect_int_int_str(batches: &[RecordBatch]) -> Vec<(i32, i32, String)> {
    let mut rows = Vec::new();
    for batch in batches {
        let col0 = batch
            .column(0)
            .as_any()
            .downcast_ref::<Int32Array>()
            .unwrap();
        let col1 = batch
            .column(1)
            .as_any()
            .downcast_ref::<Int32Array>()
            .unwrap();
        let col2 = batch.column(2);
        for i in 0..batch.num_rows() {
            rows.push((
                col0.value(i),
                col1.value(i),
                string_value(col2.as_ref(), i).to_string(),
            ));
        }
    }
    rows.sort_by_key(|r| (r.0, r.1));
    rows
}

/// Collect (i32, String) rows from batches, sorted by col0.
#[allow(dead_code)]
pub fn collect_int_str(batches: &[RecordBatch]) -> Vec<(i32, String)> {
    let mut rows = Vec::new();
    for batch in batches {
        let col0 = batch
            .column(0)
            .as_any()
            .downcast_ref::<Int32Array>()
            .unwrap();
        let col1 = batch.column(1);
        for i in 0..batch.num_rows() {
            rows.push((col0.value(i), string_value(col1.as_ref(), i).to_string()));
        }
    }
    rows.sort_by_key(|r| r.0);
    rows
}

/// Collect (i32, i32, i32) rows from batches, sorted by (col2, col0).
#[allow(dead_code)]
pub fn collect_three_ints(batches: &[RecordBatch]) -> Vec<(i32, i32, i32)> {
    let mut rows = Vec::new();
    for batch in batches {
        let col0 = batch
            .column(0)
            .as_any()
            .downcast_ref::<Int32Array>()
            .unwrap();
        let col1 = batch
            .column(1)
            .as_any()
            .downcast_ref::<Int32Array>()
            .unwrap();
        let col2 = batch
            .column(2)
            .as_any()
            .downcast_ref::<Int32Array>()
            .unwrap();
        for i in 0..batch.num_rows() {
            rows.push((col0.value(i), col1.value(i), col2.value(i)));
        }
    }
    rows.sort_by_key(|r| (r.2, r.0));
    rows
}

/// Collect (i32, String, i32) rows from batches, sorted by col0.
#[allow(dead_code)]
pub fn collect_int_str_int(batches: &[RecordBatch]) -> Vec<(i32, String, i32)> {
    let mut rows = Vec::new();
    for batch in batches {
        let col0 = batch
            .column(0)
            .as_any()
            .downcast_ref::<Int32Array>()
            .unwrap();
        let col1 = batch.column(1);
        let col2 = batch
            .column(2)
            .as_any()
            .downcast_ref::<Int32Array>()
            .unwrap();
        for i in 0..batch.num_rows() {
            rows.push((
                col0.value(i),
                string_value(col1.as_ref(), i).to_string(),
                col2.value(i),
            ));
        }
    }
    rows.sort_by_key(|r| r.0);
    rows
}

/// Query a 3-column (i32, String, i32) table and return sorted rows.
#[allow(dead_code)]
pub async fn query_int_str_int(sql_context: &SQLContext, sql: &str) -> Vec<(i32, String, i32)> {
    let batches = sql_context.sql(sql).await.unwrap().collect().await.unwrap();
    collect_int_str_int(&batches)
}