Skip to content

Commit 4f9068d

Browse files
authored
Merge pull request #2 from sdf-jkl/TLP
Add TLP Where Oracle
2 parents ba1f8ac + 1c67ae6 commit 4f9068d

8 files changed

Lines changed: 641 additions & 29 deletions

File tree

datafusion-fuzzer.toml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,7 @@ max_expr_level = 3
3131
max_table_count = 3
3232
max_insert_per_table = 20
3333

34-
# Supported oracles: NoCrash, NestedQueries.
34+
# Supported oracles: NoCrash, NestedQueries, TlpWhere.
3535
# Randomly select one oracle from the configured set for each query.
3636
oracles = ["NoCrash"]
37-
# oracles = ["NoCrash", "NestedQueries"]
37+
# oracles = ["NoCrash", "NestedQueries", "TlpWhere"]

src/common/util.rs

Lines changed: 88 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,95 @@
1-
use datafusion::{prelude::Expr, sql::unparser::expr_to_sql};
1+
use datafusion::{
2+
arrow::array::RecordBatch, common::utils::get_row_at_idx, prelude::Expr, scalar::ScalarValue,
3+
sql::unparser::expr_to_sql,
4+
};
5+
use std::collections::HashMap;
26

3-
use super::Result;
7+
use super::{Result, fuzzer_err};
48

59
/// Convert a DataFusion `Expr` into a SQL string using DataFusion's unparser.
610
pub fn to_sql_string(expr: &Expr) -> Result<String> {
711
let unparsed = expr_to_sql(expr)?;
812
Ok(unparsed.to_string())
913
}
14+
15+
pub(crate) type RowMultiset = HashMap<Vec<ScalarValue>, usize>;
16+
17+
pub(crate) fn batches_to_row_multiset(batches: &[RecordBatch]) -> Result<RowMultiset> {
18+
let mut multiset: RowMultiset = HashMap::new();
19+
let mut expected_num_cols: Option<usize> = None;
20+
21+
for batch in batches {
22+
if let Some(expected) = expected_num_cols {
23+
if batch.num_columns() != expected {
24+
return Err(fuzzer_err(&format!(
25+
"Mismatched column count across batches: expected {}, got {}",
26+
expected,
27+
batch.num_columns()
28+
)));
29+
}
30+
} else {
31+
expected_num_cols = Some(batch.num_columns());
32+
}
33+
34+
for row_idx in 0..batch.num_rows() {
35+
let mut row_key = get_row_at_idx(batch.columns(), row_idx)
36+
.map_err(|e| fuzzer_err(&format!("Failed to extract row {}: {}", row_idx, e)))?;
37+
row_key.iter_mut().for_each(|v| *v = v.clone().compacted());
38+
*multiset.entry(row_key).or_insert(0) += 1;
39+
}
40+
}
41+
42+
Ok(multiset)
43+
}
44+
45+
pub(crate) fn format_row_multiset_diff(left: &RowMultiset, right: &RowMultiset) -> String {
46+
let mut lines = Vec::new();
47+
for (row, left_count) in left {
48+
let right_count = right.get(row).copied().unwrap_or(0);
49+
if *left_count != right_count {
50+
lines.push(format!(
51+
"row={:?}, left_count={}, right_count={}",
52+
row, left_count, right_count
53+
));
54+
}
55+
}
56+
for (row, right_count) in right {
57+
if !left.contains_key(row) {
58+
lines.push(format!(
59+
"row={:?}, left_count=0, right_count={}",
60+
row, right_count
61+
));
62+
}
63+
}
64+
65+
lines.sort();
66+
let preview = lines.into_iter().take(20).collect::<Vec<_>>();
67+
if preview.is_empty() {
68+
"no row differences".to_string()
69+
} else {
70+
preview.join("\n")
71+
}
72+
}
73+
74+
pub(crate) fn count_total_rows(batches: &[RecordBatch]) -> usize {
75+
batches.iter().map(RecordBatch::num_rows).sum()
76+
}
77+
78+
pub(crate) fn validate_batches_value_equivalence(
79+
left_batches: &[RecordBatch],
80+
right_batches: &[RecordBatch],
81+
oracle_name: &str,
82+
) -> Result<()> {
83+
let left_multiset = batches_to_row_multiset(left_batches)?;
84+
let right_multiset = batches_to_row_multiset(right_batches)?;
85+
86+
if left_multiset != right_multiset {
87+
return Err(fuzzer_err(&format!(
88+
"{} value equivalence violated:\n{}",
89+
oracle_name,
90+
format_row_multiset_diff(&left_multiset, &right_multiset)
91+
)));
92+
}
93+
94+
Ok(())
95+
}

src/fuzz_context/runner_config.rs

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -169,14 +169,18 @@ max_row_count = 100
169169
max_expr_level = 3
170170
max_table_count = 3
171171
max_insert_per_table = 20
172-
oracles = ["NoCrash", "NestedQueries"]
172+
oracles = ["NoCrash", "NestedQueries", "TlpWhere"]
173173
"#,
174174
)
175175
.unwrap();
176176

177177
assert_eq!(
178178
config.oracles,
179-
vec![ConfiguredOracle::NoCrash, ConfiguredOracle::NestedQueries]
179+
vec![
180+
ConfiguredOracle::NoCrash,
181+
ConfiguredOracle::NestedQueries,
182+
ConfiguredOracle::TlpWhere
183+
]
180184
);
181185
}
182186

src/oracle/mod.rs

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
// Oracle module - provides testing oracles for query consistency and correctness
22

3+
pub(crate) mod oracle_common;
34
pub mod oracle_impl_nested_queries;
45
pub mod oracle_impl_no_crash;
6+
pub mod oracle_impl_tlp_where;
57
pub mod oracle_trait;
68

79
use std::sync::Arc;
@@ -13,6 +15,7 @@ use crate::fuzz_context::GlobalContext;
1315
// Re-export main types and traits
1416
pub use oracle_impl_nested_queries::NestedQueriesOracle;
1517
pub use oracle_impl_no_crash::NoCrashOracle;
18+
pub use oracle_impl_tlp_where::TlpWhereOracle;
1619
pub use oracle_trait::{Oracle, QueryContext, QueryExecutionResult};
1720

1821
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
@@ -21,13 +24,16 @@ pub enum ConfiguredOracle {
2124
NoCrash,
2225
#[serde(rename = "NestedQueries", alias = "NestedQueriesOracle")]
2326
NestedQueries,
27+
#[serde(rename = "TlpWhere", alias = "TlpWhereOracle")]
28+
TlpWhere,
2429
}
2530

2631
impl ConfiguredOracle {
2732
pub fn build(self, seed: u64, ctx: Arc<GlobalContext>) -> Box<dyn Oracle + Send> {
2833
match self {
2934
Self::NoCrash => Box::new(NoCrashOracle::new(seed, ctx)),
3035
Self::NestedQueries => Box::new(NestedQueriesOracle::new(seed, ctx)),
36+
Self::TlpWhere => Box::new(TlpWhereOracle::new(seed, ctx)),
3137
}
3238
}
3339
}

src/oracle/oracle_common.rs

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,27 @@
1+
use crate::common::{Result, fuzzer_err, util};
2+
use crate::oracle::QueryExecutionResult;
3+
4+
pub(crate) fn validate_value_equivalence(
5+
results: &[QueryExecutionResult],
6+
left_idx: usize,
7+
right_idx: usize,
8+
oracle_name: &str,
9+
) -> Result<()> {
10+
let left_result = results
11+
.get(left_idx)
12+
.ok_or_else(|| fuzzer_err(&format!("Missing result at index {}", left_idx)))?;
13+
let right_result = results
14+
.get(right_idx)
15+
.ok_or_else(|| fuzzer_err(&format!("Missing result at index {}", right_idx)))?;
16+
17+
let left_batches = left_result
18+
.result
19+
.as_ref()
20+
.map_err(|e| fuzzer_err(&e.to_string()))?;
21+
let right_batches = right_result
22+
.result
23+
.as_ref()
24+
.map_err(|e| fuzzer_err(&e.to_string()))?;
25+
26+
util::validate_batches_value_equivalence(left_batches, right_batches, oracle_name)
27+
}

0 commit comments

Comments
 (0)