datafusion_federation/optimizer/
mod.rs1mod scan_result;
2
3use std::any::Any;
4use std::sync::Arc;
5
6use datafusion::{
7 common::not_impl_err,
8 common::tree_node::{Transformed, TreeNode, TreeNodeRecursion},
9 datasource::source_as_provider,
10 error::Result,
11 logical_expr::{Expr, Extension, LogicalPlan, Projection, TableScan, TableSource},
12 optimizer::optimizer::{Optimizer, OptimizerConfig, OptimizerRule},
13};
14
15use crate::{
16 FederatedTableProviderAdaptor, FederatedTableSource, FederationProvider, FederationProviderRef,
17};
18
19use scan_result::ScanResult;
20
21#[derive(Default, Debug)]
27pub struct FederationOptimizerRule {}
28
29impl OptimizerRule for FederationOptimizerRule {
30 fn rewrite(
36 &self,
37 plan: LogicalPlan,
38 config: &dyn OptimizerConfig,
39 ) -> Result<Transformed<LogicalPlan>> {
40 match self.optimize_plan_recursively(&plan, true, config)? {
41 (Some(optimized_plan), _) => Ok(Transformed::yes(optimized_plan)),
42 (None, _) => Ok(Transformed::no(plan)),
43 }
44 }
45
46 fn supports_rewrite(&self) -> bool {
48 true
49 }
50
51 fn name(&self) -> &str {
53 "federation_optimizer_rule"
54 }
55}
56
57impl FederationOptimizerRule {
58 pub fn new() -> Self {
60 Self::default()
61 }
62
63 fn scan_plan_recursively(&self, plan: &LogicalPlan) -> Result<ScanResult> {
65 let mut sole_provider: ScanResult = ScanResult::None;
66
67 plan.apply(&mut |p: &LogicalPlan| -> Result<TreeNodeRecursion> {
68 let exprs_provider = self.scan_plan_exprs(p)?;
69 sole_provider.merge(exprs_provider);
70
71 if sole_provider.is_ambiguous() {
72 return Ok(TreeNodeRecursion::Stop);
73 }
74
75 let sub_provider = get_leaf_provider(p)?;
76 sole_provider.add(sub_provider);
77
78 Ok(sole_provider.check_recursion())
79 })?;
80
81 Ok(sole_provider)
82 }
83
84 fn scan_plan_exprs(&self, plan: &LogicalPlan) -> Result<ScanResult> {
86 let mut sole_provider: ScanResult = ScanResult::None;
87
88 let exprs = plan.expressions();
89 for expr in &exprs {
90 let expr_result = self.scan_expr_recursively(expr)?;
91 sole_provider.merge(expr_result);
92
93 if sole_provider.is_ambiguous() {
94 return Ok(sole_provider);
95 }
96 }
97
98 Ok(sole_provider)
99 }
100
101 fn scan_expr_recursively(&self, expr: &Expr) -> Result<ScanResult> {
103 let mut sole_provider: ScanResult = ScanResult::None;
104
105 expr.apply(&mut |e: &Expr| -> Result<TreeNodeRecursion> {
106 match e {
108 Expr::ScalarSubquery(ref subquery) => {
109 let plan_result = self.scan_plan_recursively(&subquery.subquery)?;
110
111 sole_provider.merge(plan_result);
112 Ok(sole_provider.check_recursion())
113 }
114 Expr::InSubquery(_) => not_impl_err!("InSubquery"),
115 Expr::OuterReferenceColumn(..) => {
116 sole_provider = ScanResult::Ambiguous;
120 Ok(TreeNodeRecursion::Stop)
121 }
122 _ => Ok(TreeNodeRecursion::Continue),
123 }
124 })?;
125
126 Ok(sole_provider)
127 }
128
129 fn optimize_plan_recursively(
136 &self,
137 plan: &LogicalPlan,
138 is_root: bool,
139 _config: &dyn OptimizerConfig,
140 ) -> Result<(Option<LogicalPlan>, ScanResult)> {
141 let mut sole_provider: ScanResult = ScanResult::None;
142
143 if let LogicalPlan::Extension(Extension { ref node }) = plan {
144 if node.name() == "Federated" {
145 return Ok((None, ScanResult::Ambiguous));
147 }
148 }
149
150 let leaf_provider = get_leaf_provider(plan)?;
152
153 let exprs_result = self.scan_plan_exprs(plan)?;
155 let optimize_expressions = exprs_result.is_some();
156
157 if leaf_provider.is_some() && (exprs_result.is_none() || exprs_result == leaf_provider) {
159 return Ok((None, leaf_provider.into()));
160 }
161 sole_provider.add(leaf_provider);
163 sole_provider.merge(exprs_result);
164
165 let inputs = plan.inputs();
166 if inputs.is_empty() && sole_provider.is_none() {
168 return Ok((None, ScanResult::None));
169 }
170
171 let input_results = inputs
173 .iter()
174 .map(|i| self.optimize_plan_recursively(i, false, _config))
175 .collect::<Result<Vec<_>>>()?;
176
177 input_results.iter().for_each(|(_, scan_result)| {
179 sole_provider.merge(scan_result.clone());
180 });
181
182 if sole_provider.is_none() {
183 return Ok((None, ScanResult::None));
186 }
187
188 if let ScanResult::Distinct(provider) = sole_provider {
190 if !is_root {
191 return Ok((None, ScanResult::Distinct(provider)));
193 }
194
195 if matches!(plan, LogicalPlan::Analyze(_)) {
200 } else {
202 let Some(optimizer) = provider.optimizer() else {
203 return Ok((None, ScanResult::None));
205 };
206
207 let optimized = optimizer.optimize(plan.clone(), _config, |_, _| {})?;
209 return Ok((Some(optimized), ScanResult::None));
210 }
211 }
212
213 let new_inputs = input_results
219 .into_iter()
220 .enumerate()
221 .map(|(i, (input_plan, input_result))| {
222 if let Some(federated_plan) = input_plan {
223 return Ok(federated_plan);
225 }
226
227 let original_input = (*inputs.get(i).unwrap()).clone();
228 if input_result.is_ambiguous() {
229 return Ok(original_input);
232 }
233
234 let provider = input_result.unwrap();
235 let Some(provider) = provider else {
236 return Ok(original_input);
238 };
239
240 let Some(optimizer) = provider.optimizer() else {
241 return Ok(original_input);
243 };
244
245 let wrapped = wrap_projection(original_input)?;
247 let optimized = optimizer.optimize(wrapped, _config, |_, _| {})?;
248
249 Ok(optimized)
250 })
251 .collect::<Result<Vec<_>>>()?;
252
253 let new_expressions = if optimize_expressions {
255 self.optimize_plan_exprs(plan, _config)?
256 } else {
257 plan.expressions()
258 };
259
260 let new_plan = plan.with_new_exprs(new_expressions, new_inputs)?;
262
263 Ok((Some(new_plan), ScanResult::Ambiguous))
265 }
266
267 fn optimize_plan_exprs(
269 &self,
270 plan: &LogicalPlan,
271 _config: &dyn OptimizerConfig,
272 ) -> Result<Vec<Expr>> {
273 plan.expressions()
274 .iter()
275 .map(|expr| {
276 let transformed = expr
277 .clone()
278 .transform(&|e| self.optimize_expr_recursively(e, _config))?;
279 Ok(transformed.data)
280 })
281 .collect::<Result<Vec<_>>>()
282 }
283
284 fn optimize_expr_recursively(
287 &self,
288 expr: Expr,
289 _config: &dyn OptimizerConfig,
290 ) -> Result<Transformed<Expr>> {
291 match expr {
292 Expr::ScalarSubquery(ref subquery) => {
293 let (new_subquery, _) =
295 self.optimize_plan_recursively(&subquery.subquery, true, _config)?;
296 let Some(new_subquery) = new_subquery else {
297 return Ok(Transformed::no(expr));
298 };
299 Ok(Transformed::yes(Expr::ScalarSubquery(
300 subquery.with_plan(new_subquery.into()),
301 )))
302 }
303 Expr::InSubquery(_) => not_impl_err!("InSubquery"),
304 _ => Ok(Transformed::no(expr)),
305 }
306 }
307}
308
309struct NopFederationProvider {}
312
313impl FederationProvider for NopFederationProvider {
314 fn name(&self) -> &str {
315 "nop"
316 }
317
318 fn compute_context(&self) -> Option<String> {
319 None
320 }
321
322 fn optimizer(&self) -> Option<Arc<Optimizer>> {
323 None
324 }
325}
326
327fn get_leaf_provider(plan: &LogicalPlan) -> Result<Option<FederationProviderRef>> {
328 match plan {
329 LogicalPlan::TableScan(TableScan { ref source, .. }) => {
330 let Some(federated_source) = get_table_source(source)? else {
331 return Ok(Some(Arc::new(NopFederationProvider {})));
334 };
335 let provider = federated_source.federation_provider();
336 Ok(Some(provider))
337 }
338 _ => Ok(None),
339 }
340}
341
342fn wrap_projection(plan: LogicalPlan) -> Result<LogicalPlan> {
343 match plan {
345 LogicalPlan::Projection(_) => Ok(plan),
346 _ => {
347 let expr = plan
348 .schema()
349 .columns()
350 .iter()
351 .map(|c| Expr::Column(c.clone()))
352 .collect::<Vec<Expr>>();
353 Ok(LogicalPlan::Projection(Projection::try_new(
354 expr,
355 Arc::new(plan),
356 )?))
357 }
358 }
359}
360
361pub fn get_table_source(
362 source: &Arc<dyn TableSource>,
363) -> Result<Option<Arc<dyn FederatedTableSource>>> {
364 let source = source_as_provider(source)?;
366
367 let Some(wrapper) =
369 (source.as_ref() as &dyn Any).downcast_ref::<FederatedTableProviderAdaptor>()
370 else {
371 return Ok(None);
372 };
373
374 Ok(Some(Arc::clone(&wrapper.source)))
376}