krishiv_sql/
coop_amplifiers.rs1use std::sync::Arc;
21
22use datafusion::common::Result;
23use datafusion::common::config::ConfigOptions;
24use datafusion::common::tree_node::{Transformed, TreeNode};
25use datafusion::physical_optimizer::PhysicalOptimizerRule;
26use datafusion::physical_plan::ExecutionPlan;
27use datafusion::physical_plan::coop::CooperativeExec;
28use datafusion::physical_plan::joins::{CrossJoinExec, NestedLoopJoinExec};
29use datafusion::physical_plan::unnest::UnnestExec;
30
31#[derive(Debug, Default)]
35pub struct CooperativeAmplifiers {}
36
37impl CooperativeAmplifiers {
38 pub fn new() -> Self {
39 Self {}
40 }
41}
42
43fn is_amplifier(plan: &dyn ExecutionPlan) -> bool {
44 let any = plan as &dyn std::any::Any;
46 any.downcast_ref::<CrossJoinExec>().is_some()
47 || any.downcast_ref::<NestedLoopJoinExec>().is_some()
48 || any.downcast_ref::<UnnestExec>().is_some()
52}
53
54fn is_cooperative(plan: &dyn ExecutionPlan) -> bool {
56 (plan as &dyn std::any::Any)
57 .downcast_ref::<CooperativeExec>()
58 .is_some()
59}
60
61impl PhysicalOptimizerRule for CooperativeAmplifiers {
62 fn optimize(
63 &self,
64 plan: Arc<dyn ExecutionPlan>,
65 _config: &ConfigOptions,
66 ) -> Result<Arc<dyn ExecutionPlan>> {
67 plan.transform_up(|node| {
68 if is_amplifier(node.as_ref()) {
69 return Ok(Transformed::yes(
70 Arc::new(CooperativeExec::new(node)) as Arc<dyn ExecutionPlan>
71 ));
72 }
73 if is_cooperative(node.as_ref())
80 && node
81 .children()
82 .first()
83 .is_some_and(|child| is_cooperative(child.as_ref()))
84 && let Some(child) = node.children().first()
85 {
86 return Ok(Transformed::yes(Arc::clone(child)));
87 }
88 Ok(Transformed::no(node))
89 })
90 .map(|t| t.data)
91 }
92
93 fn name(&self) -> &str {
94 "CooperativeAmplifiers"
95 }
96
97 fn schema_check(&self) -> bool {
98 true
100 }
101}
102
103#[cfg(test)]
104#[allow(clippy::unwrap_used, clippy::expect_used)]
105mod tests {
106 use super::*;
107 use datafusion::arrow::datatypes::{DataType, Field, Schema};
108 use datafusion::datasource::memory::MemorySourceConfig;
109 use datafusion::physical_plan::displayable;
110
111 fn leaf() -> Arc<dyn ExecutionPlan> {
112 let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)]));
113 MemorySourceConfig::try_new_exec(&[vec![]], schema, None).unwrap()
114 }
115
116 fn optimize(plan: Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
117 CooperativeAmplifiers::new()
118 .optimize(plan, &ConfigOptions::default())
119 .unwrap()
120 }
121
122 fn rendered(plan: &Arc<dyn ExecutionPlan>) -> String {
123 format!("{}", displayable(plan.as_ref()).indent(false))
124 }
125
126 #[test]
128 fn a_cross_join_is_wrapped() {
129 let join = Arc::new(CrossJoinExec::new(leaf(), leaf())) as Arc<dyn ExecutionPlan>;
130 assert_eq!(rendered(&join).matches("Cooperative").count(), 0);
131 let out = optimize(join);
132 assert_eq!(
133 rendered(&out).matches("Cooperative").count(),
134 1,
135 "cross join not wrapped:\n{}",
136 rendered(&out)
137 );
138 }
139
140 #[test]
147 fn optimizing_twice_is_a_no_op() {
148 let join = Arc::new(CrossJoinExec::new(leaf(), leaf())) as Arc<dyn ExecutionPlan>;
149 let once = optimize(join);
150 let twice = optimize(Arc::clone(&once));
151 assert_eq!(
152 rendered(&once),
153 rendered(&twice),
154 "the rule is not idempotent; a second pass added a layer"
155 );
156 assert_eq!(
157 rendered(&twice).matches("Cooperative").count(),
158 1,
159 "expected exactly one wrapper:\n{}",
160 rendered(&twice)
161 );
162 }
163
164 #[test]
168 fn a_plan_without_an_amplifier_is_left_alone() {
169 let plan = leaf();
170 let out = optimize(Arc::clone(&plan));
171 assert_eq!(rendered(&plan), rendered(&out));
172 assert_eq!(rendered(&out).matches("Cooperative").count(), 0);
173 }
174
175 async fn physical(sql: &str) -> Arc<dyn ExecutionPlan> {
181 use datafusion::arrow::array::Int64Array;
182 use datafusion::arrow::record_batch::RecordBatch;
183 use datafusion::datasource::MemTable;
184 use datafusion::prelude::SessionContext;
185
186 let ctx = SessionContext::new();
187 let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)]));
188 for name in ["t1", "t2"] {
189 let batch = RecordBatch::try_new(
190 Arc::clone(&schema),
191 vec![Arc::new(Int64Array::from(vec![1i64, 2, 3]))],
192 )
193 .unwrap();
194 let table = MemTable::try_new(Arc::clone(&schema), vec![vec![batch]]).unwrap();
195 ctx.register_table(name, Arc::new(table)).unwrap();
196 }
197 let logical = ctx.sql(sql).await.unwrap().into_optimized_plan().unwrap();
198 ctx.state().create_physical_plan(&logical).await.unwrap()
199 }
200
201 #[tokio::test]
207 async fn a_nested_loop_join_is_wrapped() {
208 let plan = physical("SELECT t1.a FROM t1, t2 WHERE t1.a < t2.a").await;
209 assert!(
210 rendered(&plan).contains("NestedLoopJoin"),
211 "fixture stopped producing a nested-loop join:\n{}",
212 rendered(&plan)
213 );
214 let out = optimize(plan);
215 assert!(
216 rendered(&out).contains("Cooperative"),
217 "a nested-loop join must be made preemptible:\n{}",
218 rendered(&out)
219 );
220 }
221
222 #[tokio::test]
228 async fn an_unnest_is_wrapped() {
229 let plan = physical("SELECT unnest([1, 2, 3]) AS u FROM t1").await;
230 assert!(
231 rendered(&plan).contains("Unnest"),
232 "fixture stopped producing an unnest:\n{}",
233 rendered(&plan)
234 );
235 let out = optimize(plan);
236 assert!(
237 rendered(&out).contains("Cooperative"),
238 "an unnest must be made preemptible:\n{}",
239 rendered(&out)
240 );
241 }
242
243 #[tokio::test]
249 async fn every_amplifier_is_idempotent_under_a_second_pass() {
250 for sql in [
251 "SELECT t1.a FROM t1, t2 WHERE t1.a < t2.a",
252 "SELECT unnest([1, 2, 3]) AS u FROM t1",
253 ] {
254 let once = optimize(physical(sql).await);
255 let twice = optimize(Arc::clone(&once));
256 assert_eq!(
257 rendered(&once),
258 rendered(&twice),
259 "a second pass changed the plan for:\n{sql}"
260 );
261 }
262 }
263}