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()
}
#[allow(dead_code)]
pub async fn exec(sql_context: &SQLContext, s: &str) {
sql_context.sql(s).await.unwrap().collect().await.unwrap();
}
#[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}'"
);
}
#[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
}
#[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
}
#[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
}
#[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
}
#[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)
}