1#[allow(clippy::disallowed_types)] use std::collections::{HashMap, HashSet};
9use std::fmt;
10use std::sync::Arc;
11
12use datafusion::common::Result;
13use datafusion::logical_expr::logical_plan::LogicalPlan;
14use datafusion::logical_expr::{Extension, Join, TableScan, UserDefinedLogicalNodeCore};
15use datafusion_common::tree_node::Transformed;
16use datafusion_optimizer::optimizer::{ApplyOrder, OptimizerConfig, OptimizerRule};
17
18use crate::datafusion::lookup_join::{
19 JoinKeyPair, LookupJoinNode, LookupJoinType, LookupTableMetadata,
20};
21use crate::planner::LookupTableInfo;
22
23#[derive(Debug)]
26pub struct LookupJoinRewriteRule {
27 lookup_tables: HashMap<String, LookupTableInfo>,
29}
30
31impl LookupJoinRewriteRule {
32 #[must_use]
34 pub fn new(lookup_tables: HashMap<String, LookupTableInfo>) -> Self {
35 Self { lookup_tables }
36 }
37
38 fn detect_lookup_side(&self, join: &Join) -> Option<(bool, String)> {
41 if let Some(name) = scan_table_name(&join.right) {
43 if self.lookup_tables.contains_key(&name) {
44 return Some((true, name));
45 }
46 }
47 if let Some(name) = scan_table_name(&join.left) {
49 if self.lookup_tables.contains_key(&name) {
50 return Some((false, name));
51 }
52 }
53 None
54 }
55}
56
57impl OptimizerRule for LookupJoinRewriteRule {
58 fn name(&self) -> &'static str {
59 "lookup_join_rewrite"
60 }
61
62 fn apply_order(&self) -> Option<ApplyOrder> {
63 Some(ApplyOrder::BottomUp)
64 }
65
66 fn rewrite(
67 &self,
68 plan: LogicalPlan,
69 _config: &dyn OptimizerConfig,
70 ) -> Result<Transformed<LogicalPlan>> {
71 let LogicalPlan::Join(join) = &plan else {
72 return Ok(Transformed::no(plan));
73 };
74
75 let Some((lookup_is_right, table_name)) = self.detect_lookup_side(join) else {
76 return Ok(Transformed::no(plan));
77 };
78
79 let info = &self.lookup_tables[&table_name];
80
81 let (stream_plan, lookup_plan) = if lookup_is_right {
83 (join.left.as_ref(), join.right.as_ref())
84 } else {
85 (join.right.as_ref(), join.left.as_ref())
86 };
87
88 let stream_alias = scan_table_name_and_alias(stream_plan).and_then(|(_, a)| a);
90 let lookup_alias = scan_table_name_and_alias(lookup_plan).and_then(|(_, a)| a);
91
92 let lookup_schema = lookup_plan.schema().clone();
93
94 let join_keys: Vec<JoinKeyPair> = join
96 .on
97 .iter()
98 .map(|(left_expr, right_expr)| {
99 let lookup_expr = if lookup_is_right {
100 right_expr
101 } else {
102 left_expr
103 };
104 let stream_expr = if lookup_is_right {
105 left_expr
106 } else {
107 right_expr
108 };
109 let lookup_column = match lookup_expr {
110 datafusion::logical_expr::Expr::Column(col) => col.name.clone(),
111 other => other.to_string(),
112 };
113 JoinKeyPair {
114 stream_expr: stream_expr.clone(),
115 lookup_column,
116 }
117 })
118 .collect();
119
120 let join_type = match join.join_type {
122 datafusion::logical_expr::JoinType::Inner => LookupJoinType::Inner,
123 datafusion::logical_expr::JoinType::Left if lookup_is_right => {
124 LookupJoinType::LeftOuter
125 }
126 datafusion::logical_expr::JoinType::Right if !lookup_is_right => {
127 LookupJoinType::LeftOuter
128 }
129 _ => return Ok(Transformed::no(plan)),
130 };
131
132 let required_columns: HashSet<String> = lookup_schema
134 .fields()
135 .iter()
136 .map(|f| f.name().clone())
137 .collect();
138
139 let stream_schema = stream_plan.schema();
141 let output_schema = Arc::new(stream_schema.join(lookup_schema.as_ref())?);
142
143 let metadata = LookupTableMetadata {
144 connector: info.properties.connector.to_string(),
145 strategy: info.properties.strategy.to_string(),
146 pushdown_mode: info.properties.pushdown_mode.to_string(),
147 primary_key: info.primary_key.clone(),
148 };
149
150 let node = LookupJoinNode::new(
151 stream_plan.clone(),
152 table_name,
153 lookup_schema,
154 join_keys,
155 join_type,
156 vec![], required_columns,
158 output_schema,
159 metadata,
160 )
161 .with_aliases(lookup_alias, stream_alias);
162
163 Ok(Transformed::yes(LogicalPlan::Extension(Extension {
164 node: Arc::new(node),
165 })))
166 }
167}
168
169#[derive(Debug)]
174pub struct LookupColumnPruningRule;
175
176impl OptimizerRule for LookupColumnPruningRule {
177 fn name(&self) -> &'static str {
178 "lookup_column_pruning"
179 }
180
181 fn apply_order(&self) -> Option<ApplyOrder> {
182 Some(ApplyOrder::TopDown)
183 }
184
185 fn rewrite(
186 &self,
187 plan: LogicalPlan,
188 _config: &dyn OptimizerConfig,
189 ) -> Result<Transformed<LogicalPlan>> {
190 let LogicalPlan::Extension(ext) = &plan else {
191 return Ok(Transformed::no(plan));
192 };
193
194 let Some(node) = ext.node.as_any().downcast_ref::<LookupJoinNode>() else {
195 return Ok(Transformed::no(plan));
196 };
197
198 let schema = UserDefinedLogicalNodeCore::schema(node);
203 let used: HashSet<String> = schema
204 .fields()
205 .iter()
206 .filter(|f| node.required_lookup_columns().contains(f.name()))
207 .map(|f| f.name().clone())
208 .collect();
209
210 if used == *node.required_lookup_columns() {
211 return Ok(Transformed::no(plan));
212 }
213
214 let node_inputs = UserDefinedLogicalNodeCore::inputs(node);
216 let pruned = LookupJoinNode::new(
217 node_inputs[0].clone(),
218 node.lookup_table_name().to_string(),
219 node.lookup_schema().clone(),
220 node.join_keys().to_vec(),
221 node.join_type(),
222 node.pushdown_predicates().to_vec(),
223 used,
224 schema.clone(),
225 node.metadata().clone(),
226 )
227 .with_local_predicates(node.local_predicates().to_vec())
228 .with_aliases(
229 node.lookup_alias().map(String::from),
230 node.stream_alias().map(String::from),
231 );
232
233 Ok(Transformed::yes(LogicalPlan::Extension(Extension {
234 node: Arc::new(pruned),
235 })))
236 }
237}
238
239fn scan_table_name_and_alias(plan: &LogicalPlan) -> Option<(String, Option<String>)> {
244 match plan {
245 LogicalPlan::TableScan(TableScan { table_name, .. }) => {
246 Some((table_name.table().to_string(), None))
247 }
248 LogicalPlan::SubqueryAlias(alias) => {
249 let alias_name = alias.alias.table().to_string();
250 scan_table_name_and_alias(&alias.input).map(|(base, _)| (base, Some(alias_name)))
251 }
252 _ => None,
253 }
254}
255
256fn scan_table_name(plan: &LogicalPlan) -> Option<String> {
258 scan_table_name_and_alias(plan).map(|(name, _)| name)
259}
260
261impl fmt::Display for crate::parser::lookup_table::LookupStrategy {
262 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
263 match self {
264 Self::Replicated => write!(f, "replicated"),
265 Self::OnDemand => write!(f, "on-demand"),
266 }
267 }
268}
269
270impl fmt::Display for crate::parser::lookup_table::PushdownMode {
271 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
272 match self {
273 Self::Auto => write!(f, "auto"),
274 Self::Enabled => write!(f, "enabled"),
275 Self::Disabled => write!(f, "disabled"),
276 }
277 }
278}
279
280#[cfg(test)]
281mod tests {
282 use super::*;
283 use crate::datafusion::create_session_context;
284 use crate::parser::lookup_table::{
285 LookupConnector, LookupStrategy, LookupTableProperties, PushdownMode,
286 };
287 use arrow::datatypes::{DataType, Field, Schema};
288 use datafusion::prelude::SessionContext;
289 use datafusion_common::tree_node::TreeNode;
290 use datafusion_optimizer::optimizer::OptimizerContext;
291
292 fn test_lookup_info() -> LookupTableInfo {
293 let arrow_schema = Arc::new(Schema::new(vec![
294 Field::new("id", DataType::Int32, false),
295 Field::new("name", DataType::Utf8, true),
296 ]));
297 LookupTableInfo {
298 name: "customers".to_string(),
299 columns: vec![
300 ("id".to_string(), "INT".to_string()),
301 ("name".to_string(), "VARCHAR".to_string()),
302 ],
303 primary_key: vec!["id".to_string()],
304 properties: LookupTableProperties {
305 connector: LookupConnector::External("catalog-source".into()),
306 strategy: LookupStrategy::Replicated,
307 cache_memory: None,
308 cache_ttl: None,
309 pushdown_mode: PushdownMode::Auto,
310 },
311 arrow_schema,
312 #[allow(clippy::disallowed_types)] raw_options: std::collections::HashMap::new(),
314 }
315 }
316
317 fn register_test_tables(ctx: &SessionContext) {
318 let orders_schema = Arc::new(Schema::new(vec![
319 Field::new("order_id", DataType::Int64, false),
320 Field::new("customer_id", DataType::Int64, false),
321 Field::new("amount", DataType::Float64, false),
322 ]));
323 let customers_schema = Arc::new(Schema::new(vec![
324 Field::new("id", DataType::Int64, false),
325 Field::new("name", DataType::Utf8, true),
326 ]));
327 ctx.register_batch(
328 "orders",
329 arrow::array::RecordBatch::new_empty(orders_schema),
330 )
331 .unwrap();
332 ctx.register_batch(
333 "customers",
334 arrow::array::RecordBatch::new_empty(customers_schema),
335 )
336 .unwrap();
337 }
338
339 #[tokio::test]
340 async fn test_rewrite_join_on_lookup_table() {
341 let ctx = create_session_context();
342 register_test_tables(&ctx);
343
344 let plan = ctx
345 .sql("SELECT o.order_id, c.name FROM orders o JOIN customers c ON o.customer_id = c.id")
346 .await
347 .unwrap()
348 .into_unoptimized_plan();
349
350 let mut lookup_tables = HashMap::new();
351 lookup_tables.insert("customers".to_string(), test_lookup_info());
352 let rule = LookupJoinRewriteRule::new(lookup_tables);
353
354 let transformed = plan
355 .transform_down(|p| rule.rewrite(p, &OptimizerContext::new()))
356 .unwrap();
357
358 assert!(transformed.transformed);
360 let has_lookup = format!("{:?}", transformed.data).contains("LookupJoin");
361 assert!(has_lookup, "Expected LookupJoin in plan");
362 }
363
364 #[tokio::test]
365 async fn test_non_lookup_join_not_rewritten() {
366 let ctx = create_session_context();
367 let schema_a = Arc::new(Schema::new(vec![Field::new("id", DataType::Int64, false)]));
369 let schema_b = Arc::new(Schema::new(vec![Field::new(
370 "a_id",
371 DataType::Int64,
372 false,
373 )]));
374 ctx.register_batch("a", arrow::array::RecordBatch::new_empty(schema_a))
375 .unwrap();
376 ctx.register_batch("b", arrow::array::RecordBatch::new_empty(schema_b))
377 .unwrap();
378
379 let plan = ctx
380 .sql("SELECT * FROM a JOIN b ON a.id = b.a_id")
381 .await
382 .unwrap()
383 .into_unoptimized_plan();
384
385 let rule = LookupJoinRewriteRule::new(HashMap::new());
387
388 let transformed = plan
389 .transform_down(|p| rule.rewrite(p, &OptimizerContext::new()))
390 .unwrap();
391
392 assert!(!transformed.transformed);
393 }
394
395 #[tokio::test]
396 async fn test_left_outer_produces_left_outer_type() {
397 let ctx = create_session_context();
398 register_test_tables(&ctx);
399
400 let plan = ctx
401 .sql("SELECT o.order_id, c.name FROM orders o LEFT JOIN customers c ON o.customer_id = c.id")
402 .await
403 .unwrap()
404 .into_unoptimized_plan();
405
406 let mut lookup_tables = HashMap::new();
407 lookup_tables.insert("customers".to_string(), test_lookup_info());
408 let rule = LookupJoinRewriteRule::new(lookup_tables);
409
410 let transformed = plan
411 .transform_down(|p| rule.rewrite(p, &OptimizerContext::new()))
412 .unwrap();
413
414 assert!(transformed.transformed);
415 let debug_str = format!("{:?}", transformed.data);
416 assert!(
417 debug_str.contains("LeftOuter"),
418 "Expected LeftOuter join type, got: {debug_str}"
419 );
420 }
421
422 #[test]
423 fn test_fmt_display_lookup_connector() {
424 assert_eq!(LookupConnector::Static.to_string(), "static");
425 assert_eq!(
426 LookupConnector::External("my-conn".into()).to_string(),
427 "my-conn"
428 );
429 }
430
431 #[test]
432 fn test_fmt_display_strategy() {
433 assert_eq!(LookupStrategy::Replicated.to_string(), "replicated");
434 assert_eq!(LookupStrategy::OnDemand.to_string(), "on-demand");
435 }
436
437 #[test]
438 fn test_fmt_display_pushdown_mode() {
439 assert_eq!(PushdownMode::Auto.to_string(), "auto");
440 assert_eq!(PushdownMode::Disabled.to_string(), "disabled");
441 }
442}