Skip to main content

laminar_sql/planner/
lookup_join.rs

1//! Optimizer rules that rewrite standard JOINs to `LookupJoinNode`.
2//!
3//! When a query joins a streaming source with a registered lookup table,
4//! the `LookupJoinRewriteRule` replaces the standard hash/merge join
5//! with a `LookupJoinNode` that uses the lookup source connector.
6
7#[allow(clippy::disallowed_types)] // cold path: query planning
8use 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/// Rewrites standard JOIN nodes that reference a lookup table into
24/// `LookupJoinNode` extension nodes.
25#[derive(Debug)]
26pub struct LookupJoinRewriteRule {
27    /// Registered lookup tables, keyed by name.
28    lookup_tables: HashMap<String, LookupTableInfo>,
29}
30
31impl LookupJoinRewriteRule {
32    /// Creates a new rewrite rule with the given set of registered lookup tables.
33    #[must_use]
34    pub fn new(lookup_tables: HashMap<String, LookupTableInfo>) -> Self {
35        Self { lookup_tables }
36    }
37
38    /// Detects which side of a join (if any) is a lookup table scan.
39    /// Returns `Some((lookup_side_is_right, table_name))`.
40    fn detect_lookup_side(&self, join: &Join) -> Option<(bool, String)> {
41        // Check right side
42        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        // Check left side
48        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        // Determine which side is stream and which is lookup
82        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        // Extract aliases for qualified column resolution (C7)
89        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        // Build join key pairs from the equijoin conditions
95        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        // Convert DataFusion join type to our lookup join type
121        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        // All lookup columns are required initially; pruning is done later
133        let required_columns: HashSet<String> = lookup_schema
134            .fields()
135            .iter()
136            .map(|f| f.name().clone())
137            .collect();
138
139        // Build output schema from stream + lookup
140        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![], // predicates pushed down later
157            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/// Column pruning rule for `LookupJoinNode`.
170///
171/// Narrows `required_lookup_columns` to only the columns referenced
172/// by downstream plan nodes.
173#[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        // Collect columns actually used downstream by walking the parent plan.
199        // For now, we use the node's schema to determine which lookup columns
200        // appear in the output. A full implementation would track column usage
201        // from parent nodes; this is a conservative starting point.
202        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        // Rebuild with narrowed columns
215        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
239/// Extracts the table name and optional alias from a plan node.
240///
241/// Returns `(base_table_name, alias)` — alias is the `SubqueryAlias` name
242/// if the scan is wrapped in one, `None` otherwise.
243fn 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
256/// Extracts the table name from a `TableScan` node, unwrapping aliases.
257fn 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)] // cold path: query planning
313            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        // Verify rewrite happened
359        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        // Register both as regular tables (neither is a lookup table)
368        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        // No lookup tables registered
386        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}