// Copyright 2024 RustFS Team // // 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 // // 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. #[cfg(test)] mod integration_tests { use crate::{create_fresh_db, get_global_db, instance::make_rustfsms}; use datafusion::arrow::{ array::{Array, Int64Array, StringArray}, record_batch::RecordBatch, }; use rustfs_s3select_api::{ QueryError, query::{Context, Query}, }; use s3s::dto::{ CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, JSONType, OutputSerialization, ParquetInput, ScanRange, SelectObjectContentInput, SelectObjectContentRequest, }; use std::sync::Arc; fn assert_ages_descending(output: &[RecordBatch]) { let ages: Vec = output .iter() .flat_map(|batch| { batch .column(1) .as_any() .downcast_ref::() .expect("age column should be Int64") .values() .iter() .copied() .collect::>() }) .collect(); assert_eq!(ages, vec![40, 38, 35, 32, 30, 28, 26, 25, 24, 22]); } fn assert_department_counts(output: &[RecordBatch]) { let mut counts: Vec<(&str, i64)> = output .iter() .flat_map(|batch| { let departments = batch .column(0) .as_any() .downcast_ref::() .expect("department column should be Utf8"); let counts = batch .column(1) .as_any() .downcast_ref::() .expect("count column should be Int64"); departments.iter().zip(counts.iter()).map(|(department, count)| { ( department.expect("department should not be null"), count.expect("count should not be null"), ) }) }) .collect(); counts.sort_unstable(); assert_eq!(counts, vec![("Finance", 3), ("HR", 2), ("IT", 3), ("Marketing", 2)]); } fn create_test_input(sql: &str) -> SelectObjectContentInput { SelectObjectContentInput { bucket: "test-bucket".to_string(), expected_bucket_owner: None, key: "test.csv".to_string(), sse_customer_algorithm: None, sse_customer_key: None, sse_customer_key_md5: None, request: SelectObjectContentRequest { expression: sql.to_string(), expression_type: ExpressionType::from_static("SQL"), input_serialization: InputSerialization { csv: Some(CSVInput { file_header_info: Some(FileHeaderInfo::from_static(FileHeaderInfo::USE)), ..Default::default() }), ..Default::default() }, output_serialization: OutputSerialization { csv: Some(CSVOutput::default()), ..Default::default() }, request_progress: None, scan_range: None, }, } } /// Build a `SelectObjectContentInput` targeting a JSON DOCUMENT file. /// Uses `JSONType::DOCUMENT`, which keeps the document-style validation path. fn create_test_json_input(sql: &str) -> SelectObjectContentInput { SelectObjectContentInput { bucket: "test-bucket".to_string(), expected_bucket_owner: None, key: "test.json".to_string(), sse_customer_algorithm: None, sse_customer_key: None, sse_customer_key_md5: None, request: SelectObjectContentRequest { expression: sql.to_string(), expression_type: ExpressionType::from_static("SQL"), input_serialization: InputSerialization { json: Some(JSONInput { type_: Some(JSONType::from_static(JSONType::DOCUMENT)), }), ..Default::default() }, output_serialization: OutputSerialization { json: Some(JSONOutput::default()), ..Default::default() }, request_progress: None, scan_range: None, }, } } fn create_test_json_lines_input(sql: &str) -> SelectObjectContentInput { SelectObjectContentInput { bucket: "test-bucket".to_string(), expected_bucket_owner: None, key: "test.json".to_string(), sse_customer_algorithm: None, sse_customer_key: None, sse_customer_key_md5: None, request: SelectObjectContentRequest { expression: sql.to_string(), expression_type: ExpressionType::from_static("SQL"), input_serialization: InputSerialization { json: Some(JSONInput { type_: Some(JSONType::from_static(JSONType::LINES)), }), ..Default::default() }, output_serialization: OutputSerialization { json: Some(JSONOutput::default()), ..Default::default() }, request_progress: None, scan_range: None, }, } } fn create_test_parquet_input(sql: &str) -> SelectObjectContentInput { SelectObjectContentInput { bucket: "test-bucket".to_string(), expected_bucket_owner: None, key: "test.parquet".to_string(), sse_customer_algorithm: None, sse_customer_key: None, sse_customer_key_md5: None, request: SelectObjectContentRequest { expression: sql.to_string(), expression_type: ExpressionType::from_static("SQL"), input_serialization: InputSerialization { parquet: Some(ParquetInput {}), ..Default::default() }, output_serialization: OutputSerialization { json: Some(JSONOutput::default()), ..Default::default() }, request_progress: None, scan_range: None, }, } } #[tokio::test] async fn test_database_creation() { let input = create_test_input("SELECT * FROM S3Object"); let result = make_rustfsms(Arc::new(input), true).await; assert!(result.is_ok()); } #[tokio::test] async fn test_global_db_creation() { let input = create_test_input("SELECT * FROM S3Object"); let result = get_global_db(input.clone(), true).await; assert!(result.is_ok()); } #[tokio::test] async fn test_fresh_db_creation() { let result = create_fresh_db().await; assert!(result.is_ok()); } #[tokio::test] async fn test_simple_select_query() { let sql = "SELECT * FROM S3Object"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.expect("create csv test database"); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let query_handle = result.unwrap(); let output = query_handle.result().chunk_result().await; assert!(output.is_ok()); } #[tokio::test] async fn test_csv_values_remain_strings() { let sql = "SELECT salary FROM S3Object LIMIT 1"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.expect("create CSV test database"); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let batches = db .execute(&query) .await .expect("execute CSV query") .result() .chunk_result() .await .expect("collect CSV query output"); let salaries = batches .first() .expect("CSV query should return one batch") .column(0) .as_any() .downcast_ref::() .expect("CSV column should use UTF-8 string values"); assert_eq!(salaries.value(0), "05000"); } #[tokio::test] async fn test_csv_header_modes_keep_positional_string_columns() { for (header, expected) in [ (FileHeaderInfo::IGNORE, ["05000", "6000"]), (FileHeaderInfo::NONE, ["salary", "05000"]), ] { let sql = "SELECT _5 FROM S3Object LIMIT 2"; let mut input = create_test_input(sql); input .request .input_serialization .csv .as_mut() .expect("CSV input should be configured") .file_header_info = Some(FileHeaderInfo::from_static(header)); let db = get_global_db(input.clone(), true).await.expect("create CSV test database"); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let batches = db .execute(&query) .await .expect("execute positional CSV query") .result() .chunk_result() .await .expect("collect positional CSV query output"); let values = batches .iter() .flat_map(|batch| { let column = batch .column(0) .as_any() .downcast_ref::() .expect("positional CSV column should use UTF-8 strings"); (0..column.len()).map(|row| column.value(row).to_string()).collect::>() }) .collect::>(); assert_eq!( values, expected.map(|value| value.to_string()), "unexpected values for FileHeaderInfo={header}" ); } } #[tokio::test] async fn test_csv_numeric_comparison_with_and_without_cast() { for sql in [ "SELECT name, age FROM S3Object WHERE age > 30", "SELECT name, age FROM S3Object WHERE CAST(age AS INT) > 30", ] { let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.expect("create CSV test database"); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let batches = db .execute(&query) .await .expect("execute CSV numeric comparison") .result() .chunk_result() .await .expect("collect CSV numeric comparison output"); let mut rows = Vec::new(); for batch in &batches { let names = batch .column(0) .as_any() .downcast_ref::() .expect("CSV name column should use UTF-8 strings"); let ages = batch .column(1) .as_any() .downcast_ref::() .expect("CSV age column should use UTF-8 strings"); for row in 0..batch.num_rows() { rows.push((names.value(row).to_string(), ages.value(row).to_string())); } } assert_eq!( rows, [ ("Charlie".to_string(), "35".to_string()), ("Frank".to_string(), "40".to_string()), ("Henry".to_string(), "32".to_string()), ("Jack".to_string(), "38".to_string()), ], "unexpected rows for query: {sql}" ); } } #[tokio::test] async fn test_select_with_aggregation() { let sql = "SELECT department, COUNT(*) FROM S3Object GROUP BY department"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let output = db .execute(&query) .await .expect("execute grouped CSV query") .result() .chunk_result() .await .expect("collect grouped CSV output"); assert_department_counts(&output); } #[tokio::test] async fn test_invalid_sql_syntax() { let sql = "INVALID SQL SYNTAX"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_err()); } #[tokio::test] async fn test_multi_statement_error() { let sql = "SELECT * FROM S3Object; SELECT 1;"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_err()); if let Err(QueryError::MultiStatement { num, .. }) = result { assert_eq!(num, 2); } else { panic!("Expected MultiStatement error"); } } #[tokio::test] async fn test_query_state_machine_workflow() { let sql = "SELECT * FROM S3Object"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); // Test state machine creation let state_machine = db.build_query_state_machine(query.clone()).await; assert!(state_machine.is_ok()); let state_machine = state_machine.unwrap(); // Test logical plan building let logical_plan = db.build_logical_plan(state_machine.clone()).await; assert!(logical_plan.is_ok()); // Test execution if plan exists if let Ok(Some(plan)) = logical_plan { let execution_result = db.execute_logical_plan(plan, state_machine).await; assert!(execution_result.is_ok()); } } #[tokio::test] async fn test_query_with_limit() { let sql = "SELECT * FROM S3Object LIMIT 5"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let query_handle = result.unwrap(); let output = query_handle.result().chunk_result().await.unwrap(); // Verify that we get results (exact count depends on test data) let total_rows: usize = output.iter().map(|batch| batch.num_rows()).sum(); assert!(total_rows <= 5); } #[tokio::test] async fn test_query_with_order_by() { let sql = "SELECT name, CAST(age AS BIGINT) AS age FROM S3Object ORDER BY CAST(age AS BIGINT) DESC"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let output = db .execute(&query) .await .expect("execute ordered CSV query") .result() .chunk_result() .await .expect("collect ordered CSV output"); assert_eq!(output.iter().map(|batch| batch.num_rows()).sum::(), 10); assert_ages_descending(&output); } #[tokio::test] async fn test_concurrent_queries() { let sql = "SELECT * FROM S3Object"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); // Execute multiple queries concurrently let mut handles = vec![]; for i in 0..3 { let query = Query::new( Context { input: Arc::new(input.clone()), }, format!("SELECT * FROM S3Object LIMIT {}", i + 1), ); let db_clone = db.clone(); let handle = tokio::spawn(async move { db_clone.execute(&query).await }); handles.push(handle); } // Wait for all queries to complete for handle in handles { let result = handle.await.unwrap(); assert!(result.is_ok()); } } // ────────────────────────────────────────────── // JSON-input variants of all the above tests // ────────────────────────────────────────────── #[tokio::test] async fn test_database_creation_json() { let input = create_test_json_input("SELECT * FROM S3Object"); let result = make_rustfsms(Arc::new(input), true).await; assert!(result.is_ok()); } #[tokio::test] async fn test_global_db_creation_json() { let input = create_test_json_input("SELECT * FROM S3Object"); let result = get_global_db(input.clone(), true).await; assert!(result.is_ok()); } #[tokio::test] async fn test_simple_select_query_json() { let sql = "SELECT * FROM S3Object"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let query_handle = result.unwrap(); let output = query_handle.result().chunk_result().await; assert!(output.is_ok()); } #[tokio::test] async fn test_simple_select_query_parquet() { let sql = "SELECT name, age FROM S3Object WHERE age > 25"; let input = create_test_parquet_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let output = result.unwrap().result().chunk_result().await.unwrap(); let total_rows: usize = output.iter().map(|batch| batch.num_rows()).sum(); assert_eq!(total_rows, 3); } #[tokio::test] async fn test_simple_select_query_parquet_with_scan_range_filters_row_groups() { let sql = "SELECT name, age FROM S3Object WHERE age > 25"; let mut input = create_test_parquet_input(sql); input.request.scan_range = Some(ScanRange { start: Some(0), end: Some(1), }); let db = get_global_db(input.clone(), true) .await .expect("create parquet scan range test database"); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let output = result .expect("execute parquet scan range query") .result() .chunk_result() .await .expect("collect parquet scan range query output"); let total_rows: usize = output.iter().map(|batch| batch.num_rows()).sum(); assert_eq!(total_rows, 0); } #[tokio::test] async fn test_simple_select_query_parquet_with_full_scan_range() { let sql = "SELECT * FROM S3Object"; let mut input = create_test_parquet_input(sql); input.request.scan_range = Some(ScanRange { start: Some(0), end: Some(1024), }); let db = get_global_db(input.clone(), true) .await .expect("create parquet full scan range database"); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let output = result .expect("execute parquet full scan range query") .result() .chunk_result() .await .expect("collect parquet full scan range output"); let total_rows: usize = output.iter().map(|batch| batch.num_rows()).sum(); assert_eq!(total_rows, 5); } #[tokio::test] async fn test_simple_select_query_csv_with_scan_range() { let sql = "SELECT name, age FROM S3Object LIMIT 20"; let mut input = create_test_input(sql); input.request.scan_range = Some(ScanRange { start: Some(0), end: Some(1024), }); let db = get_global_db(input.clone(), true) .await .expect("create csv scan range test database"); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let output = result .expect("execute csv scan range query") .result() .chunk_result() .await .expect("collect csv scan range output"); let total_rows: usize = output.iter().map(|batch| batch.num_rows()).sum(); assert!(total_rows > 0); } #[tokio::test] async fn test_simple_select_query_json_with_scan_range() { let sql = "SELECT name, age FROM S3Object LIMIT 20"; let mut input = create_test_json_lines_input(sql); input.request.scan_range = Some(ScanRange { start: Some(0), end: Some(1024), }); let db = get_global_db(input.clone(), true) .await .expect("create json scan range test database"); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let output = result .expect("execute json scan range query") .result() .chunk_result() .await .expect("collect json scan range output"); let total_rows: usize = output.iter().map(|batch| batch.num_rows()).sum(); assert!(total_rows > 0); } #[tokio::test] async fn test_select_with_where_clause_json() { let sql = "SELECT name, age FROM S3Object WHERE age > 30"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); } #[tokio::test] async fn test_select_with_aggregation_json() { let sql = "SELECT department, COUNT(*) FROM S3Object GROUP BY department"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let output = db .execute(&query) .await .expect("execute grouped JSON query") .result() .chunk_result() .await .expect("collect grouped JSON output"); assert_department_counts(&output); } #[tokio::test] async fn test_invalid_sql_syntax_json() { let sql = "INVALID SQL SYNTAX"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_err()); } #[tokio::test] async fn test_multi_statement_error_json() { let sql = "SELECT * FROM S3Object; SELECT 1;"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_err()); if let Err(QueryError::MultiStatement { num, .. }) = result { assert_eq!(num, 2); } else { panic!("Expected MultiStatement error"); } } #[tokio::test] async fn test_query_state_machine_workflow_json() { let sql = "SELECT * FROM S3Object"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let state_machine = db.build_query_state_machine(query.clone()).await; assert!(state_machine.is_ok()); let state_machine = state_machine.unwrap(); let logical_plan = db.build_logical_plan(state_machine.clone()).await; assert!(logical_plan.is_ok()); if let Ok(Some(plan)) = logical_plan { let execution_result = db.execute_logical_plan(plan, state_machine).await; assert!(execution_result.is_ok()); } } #[tokio::test] async fn test_query_with_limit_json() { let sql = "SELECT * FROM S3Object LIMIT 5"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; assert!(result.is_ok()); let query_handle = result.unwrap(); let output = query_handle.result().chunk_result().await.unwrap(); let total_rows: usize = output.iter().map(|batch| batch.num_rows()).sum(); assert!(total_rows <= 5); } #[tokio::test] async fn test_query_with_order_by_json() { let sql = "SELECT name, age FROM S3Object ORDER BY age DESC"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let output = db .execute(&query) .await .expect("execute ordered JSON query") .result() .chunk_result() .await .expect("collect ordered JSON output"); assert_eq!(output.iter().map(|batch| batch.num_rows()).sum::(), 10); assert_ages_descending(&output); } #[tokio::test] async fn test_concurrent_queries_json() { let sql = "SELECT * FROM S3Object"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let mut handles = vec![]; for i in 0..3 { let query = Query::new( Context { input: Arc::new(input.clone()), }, format!("SELECT * FROM S3Object LIMIT {}", i + 1), ); let db_clone = db.clone(); let handle = tokio::spawn(async move { db_clone.execute(&query).await }); handles.push(handle); } for handle in handles { let result = handle.await.unwrap(); assert!(result.is_ok()); } } }