Skip to main content

lora_compiler/
optimizer.rs

1use crate::logical::*;
2use crate::physical::*;
3use lora_analyzer::{symbols::VarId, LiteralValue, ResolvedExpr};
4use lora_ast::BinaryOp;
5use lora_store::GraphStats;
6use std::collections::BTreeSet;
7
8pub struct Optimizer;
9
10impl Default for Optimizer {
11    fn default() -> Self {
12        Self::new()
13    }
14}
15
16impl Optimizer {
17    pub fn new() -> Self {
18        Self
19    }
20
21    /// Run the rewrite pipeline. `stats` drives cost-based selection
22    /// when multiple index rewrites match the same `Filter(NodeScan)`;
23    /// pass [`GraphStats::default()`] to disable scoring (every rewrite
24    /// commits unconditionally, the pre-cost-model behaviour).
25    pub fn optimize(&mut self, mut plan: LogicalPlan, stats: &GraphStats) -> LogicalPlan {
26        self.push_filter_below_projection(&mut plan);
27        self.use_indexed_node_scans(&mut plan, stats);
28        self.use_indexed_rel_scans(&mut plan, stats);
29        self.use_index_order_for_sorts(&mut plan);
30        self.annotate_top_k_sorts(&mut plan);
31        self.remove_redundant_limit(&mut plan);
32        plan
33    }
34
35    fn push_filter_below_projection(&self, plan: &mut LogicalPlan) {
36        let len = plan.nodes.len();
37
38        for i in 0..len {
39            let input_id = match &plan.nodes[i] {
40                LogicalOp::Filter(f) => f.input,
41                _ => continue,
42            };
43
44            let Some(input) = plan.nodes.get(input_id) else {
45                continue;
46            };
47
48            if !can_push_filter_below_projection(&plan.nodes[i], input) {
49                continue;
50            }
51
52            push_filter_below_projection_at(plan, i, input_id);
53        }
54    }
55
56    /// `MATCH (n:L) WHERE n.k > $x RETURN ... ORDER BY n.k LIMIT m` with a
57    /// range index on `(L, k)`: the range scan already walks that index,
58    /// so let it emit rows in `k` order and drop the Sort. With the Sort
59    /// gone nothing blocks between the scan and the LIMIT, so the query
60    /// streams and stops after `m` rows instead of sorting every match
61    /// (keyset pagination at a cost independent of label size).
62    ///
63    /// Applies when the Sort has a single key naming the scanned
64    /// variable's range property (directly or through a projection alias)
65    /// and only Filters and non-DISTINCT Projections, which keep rows 1:1
66    /// and in order, sit between Sort and scan. The Sort node becomes a
67    /// pass-through projection. Whether the index order can really be used
68    /// is decided per execution from the bound values (see the executor);
69    /// otherwise the scan sorts its own output, so results never change.
70    fn use_index_order_for_sorts(&self, plan: &mut LogicalPlan) {
71        for i in 0..plan.nodes.len() {
72            let LogicalOp::Sort(sort) = &plan.nodes[i] else {
73                continue;
74            };
75            if sort.items.len() != 1 {
76                continue;
77            }
78            let direction = sort.items[0].direction;
79            let sort_input = sort.input;
80            let mut key_expr = sort.items[0].expr.clone();
81
82            // Walk down to the scan, resolving an alias key through the
83            // projections passed on the way.
84            let mut cursor = sort_input;
85            let scan_id = loop {
86                match &plan.nodes[cursor] {
87                    LogicalOp::Projection(p) if !p.distinct => {
88                        if let ResolvedExpr::Variable(v) = &key_expr {
89                            if let Some(item) = p.items.iter().find(|item| item.output == *v) {
90                                key_expr = item.expr.clone();
91                            }
92                        }
93                        cursor = p.input;
94                    }
95                    LogicalOp::Filter(f) => cursor = f.input,
96                    LogicalOp::NodeByPropertyRangeScan(_) => break Some(cursor),
97                    _ => break None,
98                }
99            };
100            let Some(scan_id) = scan_id else { continue };
101            let LogicalOp::NodeByPropertyRangeScan(scan) = &plan.nodes[scan_id] else {
102                continue;
103            };
104            let key_matches = matches!(
105                &key_expr,
106                ResolvedExpr::Property { expr, property }
107                    if matches!(expr.as_ref(), ResolvedExpr::Variable(v) if *v == scan.var)
108                        && *property == scan.key
109            );
110            if !key_matches || scan.input.is_some() || scan.order.is_some() {
111                continue;
112            }
113
114            if let LogicalOp::NodeByPropertyRangeScan(scan) = &mut plan.nodes[scan_id] {
115                scan.order = Some(direction);
116            }
117            plan.nodes[i] = LogicalOp::Projection(Projection {
118                input: sort_input,
119                distinct: false,
120                items: Vec::new(),
121                include_existing: true,
122            });
123        }
124    }
125
126    fn remove_redundant_limit(&self, _plan: &mut LogicalPlan) {
127        // placeholder for future rules
128    }
129
130    fn annotate_top_k_sorts(&self, plan: &mut LogicalPlan) {
131        let len = plan.nodes.len();
132        for i in 0..len {
133            let Some((input, bound)) = limit_sort_bound(&plan.nodes[i]) else {
134                continue;
135            };
136
137            if let Some(sort) = sort_op_mut(&mut plan.nodes[input]) {
138                sort.top_k = merge_top_k_bound(sort.top_k, bound);
139            }
140        }
141    }
142
143    /// Cost-driven rewrite pass. For each `Filter(NodeScan)` pattern,
144    /// collect every catalog-backed index rewrite (equality, range,
145    /// text, spatial), drop tautological candidates that select every
146    /// row, and commit the cheapest candidate when it strictly beats
147    /// the un-rewritten scan according to [`score_logical_op`].
148    fn use_indexed_node_scans(&self, plan: &mut LogicalPlan, stats: &GraphStats) {
149        let len = plan.nodes.len();
150        for i in 0..len {
151            let (input_id, predicate) = match &plan.nodes[i] {
152                LogicalOp::Filter(f) => (f.input, f.predicate.clone()),
153                _ => continue,
154            };
155            let LogicalOp::NodeScan(scan) = &plan.nodes[input_id] else {
156                continue;
157            };
158
159            let candidates = collect_index_candidates(scan, &predicate, stats);
160            let Some(best) = pick_best_candidate(&plan.nodes[input_id], candidates, stats) else {
161                continue;
162            };
163            plan.nodes[input_id] = best;
164        }
165    }
166
167    /// Rewrite `Filter(Expand(NodeScan(empty_labels), ...), pred-on-rel)`
168    /// into a relationship-targeted index scan when the predicate
169    /// matches a known shape (`r.prop CMP value`, `r.prop STARTS WITH …`,
170    /// `point.withinBBox(r.prop, …)`). The rewrite is gated on the
171    /// source NodeScan being unconstrained. The original `Filter` stays
172    /// in place so any residual predicates outside the extracted index
173    /// condition still run with normal expression semantics.
174    fn use_indexed_rel_scans(&self, plan: &mut LogicalPlan, stats: &GraphStats) {
175        let len = plan.nodes.len();
176        for i in 0..len {
177            let (filter_input, predicate) = match &plan.nodes[i] {
178                LogicalOp::Filter(f) => (f.input, f.predicate.clone()),
179                _ => continue,
180            };
181            let LogicalOp::Expand(expand) = &plan.nodes[filter_input] else {
182                continue;
183            };
184            // Variable-length expansions and inline rel-property
185            // patterns aren't handled by the index-targeted scan op.
186            if expand.range.is_some() || expand.rel_properties.is_some() {
187                continue;
188            }
189            let Some(rel_var) = expand.rel else {
190                continue;
191            };
192            // Source must be a root, unconstrained NodeScan; otherwise
193            // replacing it can bypass upstream rows or an already-bound
194            // source variable.
195            let LogicalOp::NodeScan(src_scan) = &plan.nodes[expand.input] else {
196                continue;
197            };
198            if src_scan.input.is_some() || !src_scan.labels.is_empty() {
199                continue;
200            }
201            let nodescan_input = src_scan.input;
202            let expand = expand.clone();
203
204            let candidates =
205                collect_rel_index_candidates(&expand, &predicate, rel_var, nodescan_input, stats);
206            let Some(best) = pick_best_candidate(&plan.nodes[filter_input], candidates, stats)
207            else {
208                continue;
209            };
210            // The new scan binds (src, rel, dst) and prefilters by the
211            // indexed predicate. Keep the Filter above it so unrelated
212            // conjuncts are not dropped.
213            plan.nodes[filter_input] = best;
214        }
215    }
216
217    /// Lower a logical plan by consuming it — each op's owned payload
218    /// (expressions, patterns, items) is moved into the physical op rather
219    /// than cloned. Callers should not need the logical plan after this.
220    pub fn lower_to_physical(&mut self, logical: LogicalPlan) -> PhysicalPlan {
221        let LogicalPlan { root, nodes } = logical;
222
223        let nodes = nodes.into_iter().map(lower_logical_op).collect();
224
225        PhysicalPlan { root, nodes }
226    }
227}
228
229fn can_push_filter_below_projection(filter: &LogicalOp, input: &LogicalOp) -> bool {
230    let (LogicalOp::Filter(filter), LogicalOp::Projection(proj)) = (filter, input) else {
231        return false;
232    };
233
234    if proj.distinct || proj.include_existing {
235        return false;
236    }
237
238    let output_vars: BTreeSet<VarId> = proj.items.iter().map(|item| item.output).collect();
239    let pred_vars = collect_vars(&filter.predicate);
240    !pred_vars.iter().any(|v| output_vars.contains(v))
241}
242
243fn push_filter_below_projection_at(
244    plan: &mut LogicalPlan,
245    filter_id: PlanNodeId,
246    projection_id: PlanNodeId,
247) {
248    let filter = match plan.nodes.get(filter_id).cloned() {
249        Some(LogicalOp::Filter(f)) => f,
250        _ => return,
251    };
252    let proj = match plan.nodes.get(projection_id).cloned() {
253        Some(LogicalOp::Projection(p)) => p,
254        _ => return,
255    };
256
257    plan.nodes[projection_id] = LogicalOp::Filter(Filter {
258        input: proj.input,
259        predicate: filter.predicate,
260    });
261    plan.nodes[filter_id] = LogicalOp::Projection(Projection {
262        input: projection_id,
263        distinct: proj.distinct,
264        items: proj.items,
265        include_existing: proj.include_existing,
266    });
267}
268
269fn lower_logical_op(op: LogicalOp) -> PhysicalOp {
270    match op {
271        LogicalOp::Argument(_) => PhysicalOp::Argument(ArgumentExec),
272
273        LogicalOp::NodeScan(scan) => lower_node_scan(scan),
274
275        LogicalOp::NodeByPropertyScan(scan) => {
276            PhysicalOp::NodeByPropertyScan(NodeByPropertyScanExec {
277                input: scan.input,
278                var: scan.var,
279                labels: scan.labels,
280                key: scan.key,
281                value: scan.value,
282                in_list: scan.in_list,
283            })
284        }
285
286        LogicalOp::NodeByPropertyRangeScan(scan) => {
287            PhysicalOp::NodeByPropertyRangeScan(NodeByPropertyRangeScanExec {
288                input: scan.input,
289                var: scan.var,
290                labels: scan.labels,
291                key: scan.key,
292                lo: scan.lo,
293                lo_inclusive: scan.lo_inclusive,
294                hi: scan.hi,
295                hi_inclusive: scan.hi_inclusive,
296                order: scan.order,
297            })
298        }
299
300        LogicalOp::NodeByTextScan(scan) => PhysicalOp::NodeByTextScan(NodeByTextScanExec {
301            input: scan.input,
302            var: scan.var,
303            labels: scan.labels,
304            key: scan.key,
305            predicate: scan.predicate,
306            query: scan.query,
307        }),
308
309        LogicalOp::NodeByPointScan(scan) => PhysicalOp::NodeByPointScan(NodeByPointScanExec {
310            input: scan.input,
311            var: scan.var,
312            labels: scan.labels,
313            key: scan.key,
314            predicate: scan.predicate,
315        }),
316
317        LogicalOp::RelByPropertyRangeScan(scan) => {
318            PhysicalOp::RelByPropertyRangeScan(RelByPropertyRangeScanExec {
319                input: scan.input,
320                src: scan.src,
321                rel: scan.rel,
322                dst: scan.dst,
323                types: scan.types,
324                direction: scan.direction,
325                key: scan.key,
326                lo: scan.lo,
327                lo_inclusive: scan.lo_inclusive,
328                hi: scan.hi,
329                hi_inclusive: scan.hi_inclusive,
330            })
331        }
332
333        LogicalOp::RelByTextScan(scan) => PhysicalOp::RelByTextScan(RelByTextScanExec {
334            input: scan.input,
335            src: scan.src,
336            rel: scan.rel,
337            dst: scan.dst,
338            types: scan.types,
339            direction: scan.direction,
340            key: scan.key,
341            predicate: scan.predicate,
342            query: scan.query,
343        }),
344
345        LogicalOp::RelByPointScan(scan) => PhysicalOp::RelByPointScan(RelByPointScanExec {
346            input: scan.input,
347            src: scan.src,
348            rel: scan.rel,
349            dst: scan.dst,
350            types: scan.types,
351            direction: scan.direction,
352            key: scan.key,
353            predicate: scan.predicate,
354        }),
355
356        LogicalOp::Expand(expand) => PhysicalOp::Expand(ExpandExec {
357            input: expand.input,
358            src: expand.src,
359            rel: expand.rel,
360            dst: expand.dst,
361            types: expand.types,
362            direction: expand.direction,
363            rel_properties: expand.rel_properties,
364            range: expand.range,
365        }),
366
367        LogicalOp::Filter(filter) => PhysicalOp::Filter(FilterExec {
368            input: filter.input,
369            predicate: filter.predicate,
370        }),
371
372        LogicalOp::Projection(proj) => PhysicalOp::Projection(ProjectionExec {
373            input: proj.input,
374            distinct: proj.distinct,
375            items: proj.items,
376            include_existing: proj.include_existing,
377        }),
378
379        LogicalOp::Unwind(unwind) => PhysicalOp::Unwind(UnwindExec {
380            input: unwind.input,
381            expr: unwind.expr,
382            alias: unwind.alias,
383        }),
384
385        LogicalOp::Aggregation(agg) => PhysicalOp::HashAggregation(HashAggregationExec {
386            input: agg.input,
387            group_by: agg.group_by,
388            aggregates: agg.aggregates,
389        }),
390
391        LogicalOp::Sort(sort) => PhysicalOp::Sort(SortExec {
392            input: sort.input,
393            items: sort.items,
394            top_k: sort.top_k,
395        }),
396
397        LogicalOp::Limit(limit) => PhysicalOp::Limit(LimitExec {
398            input: limit.input,
399            skip: limit.skip,
400            limit: limit.limit,
401        }),
402
403        LogicalOp::Create(create) => PhysicalOp::Create(CreateExec {
404            input: create.input,
405            pattern: create.pattern,
406        }),
407
408        LogicalOp::Merge(merge) => PhysicalOp::Merge(MergeExec {
409            input: merge.input,
410            pattern_part: merge.pattern_part,
411            actions: merge.actions,
412        }),
413
414        LogicalOp::Delete(delete) => PhysicalOp::Delete(DeleteExec {
415            input: delete.input,
416            detach: delete.detach,
417            expressions: delete.expressions,
418        }),
419
420        LogicalOp::Set(set) => PhysicalOp::Set(SetExec {
421            input: set.input,
422            items: set.items,
423        }),
424
425        LogicalOp::Remove(remove) => PhysicalOp::Remove(RemoveExec {
426            input: remove.input,
427            items: remove.items,
428        }),
429
430        LogicalOp::Foreach(foreach) => PhysicalOp::Foreach(crate::ForeachExec {
431            input: foreach.input,
432            variable: foreach.variable,
433            list: foreach.list,
434            body: foreach.body,
435        }),
436
437        LogicalOp::OptionalMatch(om) => PhysicalOp::OptionalMatch(OptionalMatchExec {
438            input: om.input,
439            inner: om.inner,
440            new_vars: om.new_vars,
441        }),
442
443        LogicalOp::PathBuild(pb) => PhysicalOp::PathBuild(PathBuildExec {
444            input: pb.input,
445            output: pb.output,
446            node_vars: pb.node_vars,
447            rel_vars: pb.rel_vars,
448            shortest_path_all: pb.shortest_path_all,
449        }),
450
451        LogicalOp::CallSubquery(cs) => PhysicalOp::CallSubquery(CallSubqueryExec {
452            input: cs.input,
453            inner: cs.inner,
454            new_vars: cs.new_vars,
455        }),
456    }
457}
458
459fn lower_node_scan(scan: NodeScan) -> PhysicalOp {
460    if scan.labels.is_empty() {
461        PhysicalOp::NodeScan(NodeScanExec {
462            input: scan.input,
463            var: scan.var,
464        })
465    } else {
466        PhysicalOp::NodeByLabelScan(NodeByLabelScanExec {
467            input: scan.input,
468            var: scan.var,
469            labels: scan.labels,
470        })
471    }
472}
473
474fn collect_vars(expr: &ResolvedExpr) -> BTreeSet<VarId> {
475    let mut vars = BTreeSet::new();
476    collect_vars_inner(expr, &mut vars);
477    vars
478}
479
480/// Build every applicable index rewrite for a `Filter(NodeScan)` site,
481/// dropping ones whose extracted predicate is trivially true (would
482/// "match everything" — see [`is_tautological_*`] helpers). Each entry
483/// is a fully-formed `LogicalOp` ready to drop into `plan.nodes[input]`.
484fn collect_index_candidates(
485    scan: &NodeScan,
486    predicate: &ResolvedExpr,
487    stats: &GraphStats,
488) -> Vec<LogicalOp> {
489    let mut out = Vec::new();
490
491    // One candidate per equality conjunct, so an indexed key wins over an
492    // unindexed one that happens to come first in the WHERE.
493    let mut seen_keys = BTreeSet::new();
494    for (key, value) in property_equalities_for_var(predicate, scan.var) {
495        if !seen_keys.insert(key.clone()) {
496            continue;
497        }
498        out.push(LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
499            input: scan.input,
500            var: scan.var,
501            labels: scan.labels.clone(),
502            key,
503            value,
504            in_list: false,
505        }));
506    }
507
508    // `var.key IN list`: one seek per distinct element. The store builds a
509    // property index on first lookup, so this costs one lookup per element;
510    // `score_logical_op` weighs it by the list length when that is known.
511    if first_simple_label(&scan.labels).is_some() {
512        for (key, list) in property_in_lists_for_var(predicate, scan.var) {
513            if !seen_keys.insert(key.clone()) {
514                continue;
515            }
516            out.push(LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
517                input: scan.input,
518                var: scan.var,
519                labels: scan.labels.clone(),
520                key,
521                value: list,
522                in_list: true,
523            }));
524        }
525    }
526
527    if let Some(bounds) = collect_range_bounds(predicate, scan.var) {
528        if !is_tautological_range(&bounds)
529            && first_simple_label(&scan.labels)
530                .is_some_and(|label| stats.has_node_range_index(label, &bounds.key))
531        {
532            out.push(LogicalOp::NodeByPropertyRangeScan(
533                NodeByPropertyRangeScan {
534                    input: scan.input,
535                    var: scan.var,
536                    labels: scan.labels.clone(),
537                    key: bounds.key,
538                    lo: bounds.lo,
539                    lo_inclusive: bounds.lo_inclusive,
540                    hi: bounds.hi,
541                    hi_inclusive: bounds.hi_inclusive,
542                    order: None,
543                },
544            ));
545        }
546    }
547
548    if let Some(candidate) = text_predicate_for_var(predicate, scan.var) {
549        if !is_tautological_text(&candidate)
550            && first_simple_label(&scan.labels)
551                .is_some_and(|label| stats.has_node_text_index(label, &candidate.key))
552        {
553            out.push(LogicalOp::NodeByTextScan(NodeByTextScan {
554                input: scan.input,
555                var: scan.var,
556                labels: scan.labels.clone(),
557                key: candidate.key,
558                predicate: candidate.predicate,
559                query: candidate.query,
560            }));
561        }
562    }
563
564    if let Some(candidate) = point_predicate_for_var(predicate, scan.var) {
565        if !is_tautological_point(&candidate)
566            && first_simple_label(&scan.labels)
567                .is_some_and(|label| stats.has_node_point_index(label, &candidate.key))
568        {
569            out.push(LogicalOp::NodeByPointScan(NodeByPointScan {
570                input: scan.input,
571                var: scan.var,
572                labels: scan.labels.clone(),
573                key: candidate.key,
574                predicate: candidate.predicate,
575            }));
576        }
577    }
578
579    out
580}
581
582/// Build every applicable rel-targeted index rewrite for a
583/// `Filter(Expand(NodeScan, …), pred-on-rel-var)` site. Mirrors
584/// [`collect_index_candidates`] for nodes; tautological predicates are
585/// dropped through the same `is_tautological_*` helpers.
586fn collect_rel_index_candidates(
587    expand: &Expand,
588    predicate: &ResolvedExpr,
589    rel_var: VarId,
590    input: Option<PlanNodeId>,
591    stats: &GraphStats,
592) -> Vec<LogicalOp> {
593    let mut out = Vec::new();
594
595    if let Some(candidate) = text_predicate_for_var(predicate, rel_var) {
596        if !is_tautological_text(&candidate)
597            && rel_types_have_index(&expand.types, |ty| {
598                stats.has_relationship_text_index(ty, &candidate.key)
599            })
600        {
601            out.push(LogicalOp::RelByTextScan(RelByTextScan {
602                input,
603                src: expand.src,
604                rel: rel_var,
605                dst: expand.dst,
606                types: expand.types.clone(),
607                direction: expand.direction,
608                key: candidate.key,
609                predicate: candidate.predicate,
610                query: candidate.query,
611            }));
612        }
613    }
614
615    if let Some(bounds) = collect_range_bounds(predicate, rel_var) {
616        if !is_tautological_range(&bounds)
617            && rel_types_have_index(&expand.types, |ty| {
618                stats.has_relationship_range_index(ty, &bounds.key)
619            })
620        {
621            out.push(LogicalOp::RelByPropertyRangeScan(RelByPropertyRangeScan {
622                input,
623                src: expand.src,
624                rel: rel_var,
625                dst: expand.dst,
626                types: expand.types.clone(),
627                direction: expand.direction,
628                key: bounds.key,
629                lo: bounds.lo,
630                lo_inclusive: bounds.lo_inclusive,
631                hi: bounds.hi,
632                hi_inclusive: bounds.hi_inclusive,
633            }));
634        }
635    }
636
637    if let Some(candidate) = point_predicate_for_var(predicate, rel_var) {
638        if !is_tautological_point(&candidate)
639            && rel_types_have_index(&expand.types, |ty| {
640                stats.has_relationship_point_index(ty, &candidate.key)
641            })
642        {
643            out.push(LogicalOp::RelByPointScan(RelByPointScan {
644                input,
645                src: expand.src,
646                rel: rel_var,
647                dst: expand.dst,
648                types: expand.types.clone(),
649                direction: expand.direction,
650                key: candidate.key,
651                predicate: candidate.predicate,
652            }));
653        }
654    }
655
656    out
657}
658
659fn rel_types_have_index<F>(types: &[String], mut has_index: F) -> bool
660where
661    F: FnMut(&str) -> bool,
662{
663    !types.is_empty() && types.iter().all(|ty| has_index(ty))
664}
665
666/// Choose the cheapest replacement for the original `Filter(NodeScan)`
667/// input among `candidates`, breaking ties by the caller-provided
668/// collection order so behaviour is deterministic. Returns
669/// `None` when no candidate is strictly cheaper than the original — in
670/// that case the caller leaves the plan unchanged.
671fn pick_best_candidate(
672    original: &LogicalOp,
673    candidates: Vec<LogicalOp>,
674    stats: &GraphStats,
675) -> Option<LogicalOp> {
676    if candidates.is_empty() {
677        return None;
678    }
679
680    let baseline = score_logical_op(original, stats);
681
682    let mut best: Option<(LogicalOp, Option<u64>)> = None;
683    for candidate in candidates {
684        let score = score_logical_op(&candidate, stats);
685        if !improves_over(score, baseline) {
686            continue;
687        }
688        let take = match &best {
689            None => true,
690            Some((_, current_best)) => is_cheaper(score, *current_best),
691        };
692        if take {
693            best = Some((candidate, score));
694        }
695    }
696
697    best.map(|(op, _)| op)
698}
699
700/// `true` when committing a candidate with `score` is at least as good
701/// as keeping the original (`baseline`). With unknown stats the
702/// optimizer keeps its pre-cost-model behaviour: a `None`/`None`
703/// comparison still commits already-collected candidates. Catalog
704/// checks happen before this point for range/text/point operators.
705fn improves_over(score: Option<u64>, baseline: Option<u64>) -> bool {
706    match (score, baseline) {
707        (Some(s), Some(b)) => s <= b,
708        (Some(_), None) => true,
709        (None, Some(_)) => false,
710        (None, None) => true,
711    }
712}
713
714fn is_cheaper(score: Option<u64>, current_best: Option<u64>) -> bool {
715    match (score, current_best) {
716        (Some(s), Some(b)) => s < b,
717        (Some(_), None) => true,
718        (None, _) => false,
719    }
720}
721
722/// Estimated row count produced by `op`. `None` means "no information"
723/// — typically because the relevant labels or property are not in the
724/// stats snapshot. Mirrors the per-operator estimates that `EXPLAIN`
725/// surfaces in [`crate::plan_tree::PlanTreeNode::estimated_rows`], so
726/// the optimizer's pick agrees with what users see in the plan tree.
727fn score_logical_op(op: &LogicalOp, stats: &GraphStats) -> Option<u64> {
728    match op {
729        LogicalOp::NodeScan(scan) => match label_estimate(&scan.labels, stats) {
730            Some(rows) => Some(rows),
731            None if scan.labels.is_empty() => Some(stats.node_count as u64),
732            None => None,
733        },
734        LogicalOp::NodeByPropertyScan(scan) => {
735            let label = first_simple_label(&scan.labels)?;
736            let per_value = stats.estimate_node_property_equality(label, &scan.key)?;
737            if !scan.in_list {
738                return Some(per_value);
739            }
740            // A literal list costs one seek per element; a parameter or
741            // other list-valued expression is assumed short.
742            let elements = match &scan.value {
743                ResolvedExpr::List(items) => items.len().max(1) as u64,
744                _ => 1,
745            };
746            Some(per_value.saturating_mul(elements))
747        }
748        LogicalOp::NodeByPropertyRangeScan(scan) => {
749            // Conservative one-third selectivity: a one-sided range
750            // typically narrows by less than half; a two-sided range
751            // narrows further. Better than full-label scan for any
752            // useful range, never worse than `label_count`.
753            let base = label_estimate(&scan.labels, stats)?;
754            let denom = match (scan.lo.is_some(), scan.hi.is_some()) {
755                (true, true) => 4,
756                _ => 3,
757            };
758            Some(base.div_ceil(denom))
759        }
760        LogicalOp::NodeByTextScan(scan) => {
761            let base = label_estimate(&scan.labels, stats)?;
762            // Prefix/suffix probes typically narrow more than CONTAINS.
763            let denom = match scan.predicate {
764                TextPredicate::StartsWith | TextPredicate::EndsWith => 4,
765                TextPredicate::Contains => 2,
766            };
767            Some(base.div_ceil(denom))
768        }
769        LogicalOp::NodeByPointScan(scan) => {
770            let base = label_estimate(&scan.labels, stats)?;
771            // Spatial probes: bbox/distance usually returns a small
772            // fraction of the labelled set.
773            Some(base.div_ceil(5))
774        }
775        LogicalOp::Filter(_) => None,
776        LogicalOp::Expand(expand) => {
777            // Used as the baseline when evaluating rel-index rewrites.
778            // Without per-edge histograms we approximate edges-of-type
779            // by `relationship_type_count`, falling back to the global
780            // relationship total when no type is named.
781            let count = if expand.types.is_empty() {
782                stats.relationship_count as u64
783            } else {
784                let mut total: u64 = 0;
785                for ty in &expand.types {
786                    total = total.saturating_add(stats.relationship_type_count(ty)?);
787                }
788                total
789            };
790            // Undirected expansion produces both orientations.
791            Some(match expand.direction {
792                lora_ast::Direction::Undirected => count.saturating_mul(2),
793                _ => count,
794            })
795        }
796        LogicalOp::RelByPropertyRangeScan(scan) => {
797            let base = rel_type_estimate(&scan.types, stats)?;
798            let denom = match (scan.lo.is_some(), scan.hi.is_some()) {
799                (true, true) => 4,
800                _ => 3,
801            };
802            let est = base.div_ceil(denom);
803            Some(match scan.direction {
804                lora_ast::Direction::Undirected => est.saturating_mul(2),
805                _ => est,
806            })
807        }
808        LogicalOp::RelByTextScan(scan) => {
809            let base = rel_type_estimate(&scan.types, stats)?;
810            let denom = match scan.predicate {
811                TextPredicate::StartsWith | TextPredicate::EndsWith => 4,
812                TextPredicate::Contains => 2,
813            };
814            let est = base.div_ceil(denom);
815            Some(match scan.direction {
816                lora_ast::Direction::Undirected => est.saturating_mul(2),
817                _ => est,
818            })
819        }
820        LogicalOp::RelByPointScan(scan) => {
821            let base = rel_type_estimate(&scan.types, stats)?;
822            let est = base.div_ceil(5);
823            Some(match scan.direction {
824                lora_ast::Direction::Undirected => est.saturating_mul(2),
825                _ => est,
826            })
827        }
828        _ => None,
829    }
830}
831
832fn rel_type_estimate(types: &[String], stats: &GraphStats) -> Option<u64> {
833    if types.is_empty() {
834        return Some(stats.relationship_count as u64);
835    }
836    let mut total: u64 = 0;
837    for ty in types {
838        total = total.saturating_add(stats.relationship_type_count(ty)?);
839    }
840    Some(total)
841}
842
843/// Return the count of nodes covered by a `labels` group, taking the
844/// first DNF disjunction's first literal. Mirrors `labels_estimate`
845/// from `lora-database/src/database/explain.rs` — keeping the two in
846/// sync ensures `EXPLAIN` and the optimizer agree.
847fn label_estimate(labels: &[Vec<String>], stats: &GraphStats) -> Option<u64> {
848    let label = first_simple_label(labels)?;
849    stats.label_count(label)
850}
851
852fn first_simple_label(labels: &[Vec<String>]) -> Option<&str> {
853    labels.first()?.first().map(String::as_str)
854}
855
856fn is_tautological_range(bounds: &RangeBounds) -> bool {
857    let lo_open = match (&bounds.lo, bounds.lo_inclusive) {
858        (None, _) => true,
859        (Some(expr), false) => matches!(
860            expr,
861            ResolvedExpr::Literal(LiteralValue::Integer(v)) if *v == i64::MIN
862        ),
863        (Some(_), true) => false,
864    };
865    let hi_open = match (&bounds.hi, bounds.hi_inclusive) {
866        (None, _) => true,
867        (Some(expr), false) => matches!(
868            expr,
869            ResolvedExpr::Literal(LiteralValue::Integer(v)) if *v == i64::MAX
870        ),
871        (Some(_), true) => false,
872    };
873    lo_open && hi_open
874}
875
876fn is_tautological_text(candidate: &TextCandidate) -> bool {
877    matches!(
878        &candidate.query,
879        ResolvedExpr::Literal(LiteralValue::String(s)) if s.is_empty()
880    )
881}
882
883fn is_tautological_point(candidate: &PointCandidate) -> bool {
884    match &candidate.predicate {
885        PointPredicate::WithinBBox {
886            lower_left,
887            upper_right,
888        } => is_world_bbox(lower_left, upper_right),
889        PointPredicate::WithinDistance { .. } => false,
890    }
891}
892
893/// `{longitude: -180, latitude: -90}::POINT` to
894/// `{longitude: 180, latitude: 90}::POINT` covers every WGS-84 point.
895/// Detecting the literal form avoids a trigram lookup that would yield
896/// every indexed row only to be re-filtered to the same set.
897fn is_world_bbox(lower_left: &ResolvedExpr, upper_right: &ResolvedExpr) -> bool {
898    /// Cypher parses `-180` as `Unary{Neg, Literal(180)}`, not as a
899    /// negative integer literal — peel one such layer so the
900    /// world-bbox detection works on natural query forms.
901    fn const_number(expr: &ResolvedExpr) -> Option<f64> {
902        match expr {
903            ResolvedExpr::Literal(LiteralValue::Float(v)) => Some(*v),
904            ResolvedExpr::Literal(LiteralValue::Integer(v)) => Some(*v as f64),
905            ResolvedExpr::Unary {
906                op: lora_ast::UnaryOp::Neg,
907                expr,
908            } => const_number(expr).map(|v| -v),
909            ResolvedExpr::Unary {
910                op: lora_ast::UnaryOp::Pos,
911                expr,
912            } => const_number(expr),
913            _ => None,
914        }
915    }
916
917    fn point_lon_lat(expr: &ResolvedExpr) -> Option<(f64, f64)> {
918        let items = point_literal_map(expr)?;
919        let mut lon: Option<f64> = None;
920        let mut lat: Option<f64> = None;
921        for (key, value) in items {
922            let n = const_number(value)?;
923            match key.as_str() {
924                "longitude" | "x" => lon = Some(n),
925                "latitude" | "y" => lat = Some(n),
926                _ => {}
927            }
928        }
929        Some((lon?, lat?))
930    }
931
932    let Some((ll_lon, ll_lat)) = point_lon_lat(lower_left) else {
933        return false;
934    };
935    let Some((ur_lon, ur_lat)) = point_lon_lat(upper_right) else {
936        return false;
937    };
938    ll_lon <= -180.0 && ll_lat <= -90.0 && ur_lon >= 180.0 && ur_lat >= 90.0
939}
940
941fn point_literal_map(expr: &ResolvedExpr) -> Option<&Vec<(String, ResolvedExpr)>> {
942    let ResolvedExpr::Function { function, args, .. } = expr else {
943        return None;
944    };
945    if function.eq_ignore_ascii_case("geo.point") && args.len() == 1 {
946        let ResolvedExpr::Map(items) = &args[0] else {
947            return None;
948        };
949        return Some(items);
950    }
951    if function.eq_ignore_ascii_case("cast.to") && args.len() == 2 {
952        let ResolvedExpr::Map(items) = &args[0] else {
953            return None;
954        };
955        let ResolvedExpr::Literal(LiteralValue::TypeName(target)) = &args[1] else {
956            return None;
957        };
958        if target.eq_ignore_ascii_case("POINT") {
959            return Some(items);
960        }
961    }
962    None
963}
964
965struct RangeBounds {
966    key: String,
967    lo: Option<ResolvedExpr>,
968    lo_inclusive: bool,
969    hi: Option<ResolvedExpr>,
970    hi_inclusive: bool,
971}
972
973/// Walk an AND-tree and collect any `var.prop CMP literal` bounds.
974/// Returns `None` if no comparison touches `var.prop` for a single
975/// property key — we don't try to combine multi-property bounds in v1.
976fn collect_range_bounds(predicate: &ResolvedExpr, var: VarId) -> Option<RangeBounds> {
977    let mut key: Option<String> = None;
978    let mut lo: Option<ResolvedExpr> = None;
979    let mut lo_inclusive = false;
980    let mut hi: Option<ResolvedExpr> = None;
981    let mut hi_inclusive = false;
982    let mut any = false;
983
984    walk_and_for_range(
985        predicate,
986        var,
987        &mut |found_key, side, value, inclusive| match side {
988            RangeSide::Lower => {
989                if key.as_deref().map(|k| k != found_key).unwrap_or(false) {
990                    return;
991                }
992                key = Some(found_key.to_string());
993                if lo
994                    .as_ref()
995                    .map(|current| lower_bound_is_tighter(&value, inclusive, current, lo_inclusive))
996                    .unwrap_or(true)
997                {
998                    lo = Some(value);
999                    lo_inclusive = inclusive;
1000                }
1001                any = true;
1002            }
1003            RangeSide::Upper => {
1004                if key.as_deref().map(|k| k != found_key).unwrap_or(false) {
1005                    return;
1006                }
1007                key = Some(found_key.to_string());
1008                if hi
1009                    .as_ref()
1010                    .map(|current| upper_bound_is_tighter(&value, inclusive, current, hi_inclusive))
1011                    .unwrap_or(true)
1012                {
1013                    hi = Some(value);
1014                    hi_inclusive = inclusive;
1015                }
1016                any = true;
1017            }
1018        },
1019    );
1020
1021    if !any {
1022        return None;
1023    }
1024    Some(RangeBounds {
1025        key: key?,
1026        lo,
1027        lo_inclusive,
1028        hi,
1029        hi_inclusive,
1030    })
1031}
1032
1033#[derive(Clone, Copy)]
1034enum RangeSide {
1035    Lower,
1036    Upper,
1037}
1038
1039fn lower_bound_is_tighter(
1040    candidate: &ResolvedExpr,
1041    candidate_inclusive: bool,
1042    current: &ResolvedExpr,
1043    current_inclusive: bool,
1044) -> bool {
1045    match compare_literal_bounds(candidate, current) {
1046        Some(std::cmp::Ordering::Greater) => true,
1047        Some(std::cmp::Ordering::Equal) => !candidate_inclusive && current_inclusive,
1048        _ => false,
1049    }
1050}
1051
1052fn upper_bound_is_tighter(
1053    candidate: &ResolvedExpr,
1054    candidate_inclusive: bool,
1055    current: &ResolvedExpr,
1056    current_inclusive: bool,
1057) -> bool {
1058    match compare_literal_bounds(candidate, current) {
1059        Some(std::cmp::Ordering::Less) => true,
1060        Some(std::cmp::Ordering::Equal) => !candidate_inclusive && current_inclusive,
1061        _ => false,
1062    }
1063}
1064
1065fn compare_literal_bounds(lhs: &ResolvedExpr, rhs: &ResolvedExpr) -> Option<std::cmp::Ordering> {
1066    match (literal_number(lhs), literal_number(rhs)) {
1067        (Some(a), Some(b)) => return a.partial_cmp(&b),
1068        (Some(_), None) | (None, Some(_)) => return None,
1069        (None, None) => {}
1070    }
1071
1072    match (lhs, rhs) {
1073        (
1074            ResolvedExpr::Literal(LiteralValue::String(a)),
1075            ResolvedExpr::Literal(LiteralValue::String(b)),
1076        ) => Some(a.cmp(b)),
1077        _ => None,
1078    }
1079}
1080
1081fn literal_number(expr: &ResolvedExpr) -> Option<f64> {
1082    match expr {
1083        ResolvedExpr::Literal(LiteralValue::Integer(v)) => Some(*v as f64),
1084        ResolvedExpr::Literal(LiteralValue::Float(v)) => Some(*v),
1085        ResolvedExpr::Unary {
1086            op: lora_ast::UnaryOp::Neg,
1087            expr,
1088        } => literal_number(expr).map(|v| -v),
1089        ResolvedExpr::Unary {
1090            op: lora_ast::UnaryOp::Pos,
1091            expr,
1092        } => literal_number(expr),
1093        _ => None,
1094    }
1095}
1096
1097fn walk_and_for_range<F>(predicate: &ResolvedExpr, var: VarId, visit: &mut F)
1098where
1099    F: FnMut(&str, RangeSide, ResolvedExpr, bool),
1100{
1101    if let ResolvedExpr::Binary {
1102        lhs,
1103        op: BinaryOp::And,
1104        rhs,
1105    } = predicate
1106    {
1107        walk_and_for_range(lhs, var, visit);
1108        walk_and_for_range(rhs, var, visit);
1109        return;
1110    }
1111
1112    let ResolvedExpr::Binary { lhs, op, rhs } = predicate else {
1113        return;
1114    };
1115
1116    let (side, inclusive) = match op {
1117        BinaryOp::Gt => (RangeSide::Lower, false),
1118        BinaryOp::Ge => (RangeSide::Lower, true),
1119        BinaryOp::Lt => (RangeSide::Upper, false),
1120        BinaryOp::Le => (RangeSide::Upper, true),
1121        _ => return,
1122    };
1123
1124    if let Some(key) = property_access_for_var(lhs, var) {
1125        if !collect_vars(rhs).contains(&var) {
1126            visit(&key, side, (**rhs).clone(), inclusive);
1127            return;
1128        }
1129    }
1130    if let Some(key) = property_access_for_var(rhs, var) {
1131        if !collect_vars(lhs).contains(&var) {
1132            // Mirror `value CMP var.prop` to `var.prop FLIPPED_CMP value`.
1133            let flipped = match side {
1134                RangeSide::Lower => RangeSide::Upper,
1135                RangeSide::Upper => RangeSide::Lower,
1136            };
1137            visit(&key, flipped, (**lhs).clone(), inclusive);
1138        }
1139    }
1140}
1141
1142struct TextCandidate {
1143    key: String,
1144    predicate: TextPredicate,
1145    query: ResolvedExpr,
1146}
1147
1148struct PointCandidate {
1149    key: String,
1150    predicate: PointPredicate,
1151}
1152
1153fn point_predicate_for_var(predicate: &ResolvedExpr, var: VarId) -> Option<PointCandidate> {
1154    if let ResolvedExpr::Binary {
1155        lhs,
1156        op: BinaryOp::And,
1157        rhs,
1158    } = predicate
1159    {
1160        return point_predicate_for_var(lhs, var).or_else(|| point_predicate_for_var(rhs, var));
1161    }
1162
1163    // geo.within_bbox(n.prop, ll, ur)
1164    if let ResolvedExpr::Function { function, args, .. } = predicate {
1165        if function.eq_ignore_ascii_case("geo.within_bbox") && args.len() == 3 {
1166            if let Some(key) = property_access_for_var(&args[0], var) {
1167                if !collect_vars(&args[1]).contains(&var) && !collect_vars(&args[2]).contains(&var)
1168                {
1169                    return Some(PointCandidate {
1170                        key,
1171                        predicate: PointPredicate::WithinBBox {
1172                            lower_left: args[1].clone(),
1173                            upper_right: args[2].clone(),
1174                        },
1175                    });
1176                }
1177            }
1178        }
1179    }
1180
1181    // point.distance(n.prop, c) OP d  (where OP is <, <=)
1182    let ResolvedExpr::Binary { lhs, op, rhs } = predicate else {
1183        return None;
1184    };
1185    let inclusive = match op {
1186        BinaryOp::Le => true,
1187        BinaryOp::Lt => false,
1188        // Symmetric form: d >= point.distance(n.prop, c)
1189        BinaryOp::Ge => true,
1190        BinaryOp::Gt => false,
1191        _ => return None,
1192    };
1193
1194    let (call_side, scalar_side) = match op {
1195        BinaryOp::Le | BinaryOp::Lt => ((**lhs).clone(), (**rhs).clone()),
1196        BinaryOp::Ge | BinaryOp::Gt => ((**rhs).clone(), (**lhs).clone()),
1197        _ => return None,
1198    };
1199
1200    let ResolvedExpr::Function { function, args, .. } = &call_side else {
1201        return None;
1202    };
1203    if !function.eq_ignore_ascii_case("geo.distance") {
1204        return None;
1205    }
1206    if args.len() != 2 {
1207        return None;
1208    }
1209    let key = property_access_for_var(&args[0], var)?;
1210    if collect_vars(&args[1]).contains(&var) || collect_vars(&scalar_side).contains(&var) {
1211        return None;
1212    }
1213    Some(PointCandidate {
1214        key,
1215        predicate: PointPredicate::WithinDistance {
1216            center: args[1].clone(),
1217            max_distance: scalar_side,
1218            inclusive,
1219        },
1220    })
1221}
1222
1223fn text_predicate_for_var(predicate: &ResolvedExpr, var: VarId) -> Option<TextCandidate> {
1224    let ResolvedExpr::Binary { lhs, op, rhs } = predicate else {
1225        return None;
1226    };
1227
1228    if matches!(op, BinaryOp::And) {
1229        return text_predicate_for_var(lhs, var).or_else(|| text_predicate_for_var(rhs, var));
1230    }
1231
1232    let kind = match op {
1233        BinaryOp::StartsWith => TextPredicate::StartsWith,
1234        BinaryOp::EndsWith => TextPredicate::EndsWith,
1235        BinaryOp::Contains => TextPredicate::Contains,
1236        _ => return None,
1237    };
1238
1239    let key = property_access_for_var(lhs, var)?;
1240    if collect_vars(rhs).contains(&var) {
1241        return None;
1242    }
1243    Some(TextCandidate {
1244        key,
1245        predicate: kind,
1246        query: (**rhs).clone(),
1247    })
1248}
1249
1250/// Every `var.key = value` conjunct of an AND-tree, in order, where
1251/// `value` does not mention `var`.
1252fn property_equalities_for_var(
1253    predicate: &ResolvedExpr,
1254    var: VarId,
1255) -> Vec<(String, ResolvedExpr)> {
1256    let mut out = Vec::new();
1257    for conjunct in and_conjuncts(predicate) {
1258        let ResolvedExpr::Binary {
1259            lhs,
1260            op: BinaryOp::Eq,
1261            rhs,
1262        } = conjunct
1263        else {
1264            continue;
1265        };
1266        if let Some(key) =
1267            property_access_for_var(lhs, var).filter(|_| !collect_vars(rhs).contains(&var))
1268        {
1269            out.push((key, (**rhs).clone()));
1270        } else if let Some(key) =
1271            property_access_for_var(rhs, var).filter(|_| !collect_vars(lhs).contains(&var))
1272        {
1273            out.push((key, (**lhs).clone()));
1274        }
1275    }
1276    out
1277}
1278
1279/// Every `var.key IN list` conjunct of an AND-tree, in order, where `list`
1280/// does not mention `var`.
1281fn property_in_lists_for_var(predicate: &ResolvedExpr, var: VarId) -> Vec<(String, ResolvedExpr)> {
1282    let mut out = Vec::new();
1283    for conjunct in and_conjuncts(predicate) {
1284        let ResolvedExpr::Binary {
1285            lhs,
1286            op: BinaryOp::In,
1287            rhs,
1288        } = conjunct
1289        else {
1290            continue;
1291        };
1292        if collect_vars(rhs).contains(&var) {
1293            continue;
1294        }
1295        if let Some(key) = property_access_for_var(lhs, var) {
1296            out.push((key, (**rhs).clone()));
1297        }
1298    }
1299    out
1300}
1301
1302/// Flatten an AND-tree into its conjuncts, left to right.
1303fn and_conjuncts(predicate: &ResolvedExpr) -> Vec<&ResolvedExpr> {
1304    let mut out = Vec::new();
1305    let mut stack = vec![predicate];
1306    while let Some(expr) = stack.pop() {
1307        match expr {
1308            ResolvedExpr::Binary {
1309                lhs,
1310                op: BinaryOp::And,
1311                rhs,
1312            } => {
1313                stack.push(rhs);
1314                stack.push(lhs);
1315            }
1316            other => out.push(other),
1317        }
1318    }
1319    out
1320}
1321
1322fn static_limit_bound(limit: &Limit) -> Option<usize> {
1323    let limit_rows = static_non_negative_usize(limit.limit.as_ref()?)?;
1324    let skip_rows = limit
1325        .skip
1326        .as_ref()
1327        .and_then(static_non_negative_usize)
1328        .unwrap_or(0);
1329    Some(skip_rows.saturating_add(limit_rows))
1330}
1331
1332fn limit_sort_bound(op: &LogicalOp) -> Option<(PlanNodeId, usize)> {
1333    let LogicalOp::Limit(limit) = op else {
1334        return None;
1335    };
1336
1337    static_limit_bound(limit).map(|bound| (limit.input, bound))
1338}
1339
1340fn sort_op_mut(op: &mut LogicalOp) -> Option<&mut Sort> {
1341    match op {
1342        LogicalOp::Sort(sort) => Some(sort),
1343        _ => None,
1344    }
1345}
1346
1347fn merge_top_k_bound(current: Option<usize>, bound: usize) -> Option<usize> {
1348    Some(current.map(|current| current.min(bound)).unwrap_or(bound))
1349}
1350
1351fn static_non_negative_usize(expr: &ResolvedExpr) -> Option<usize> {
1352    match expr {
1353        ResolvedExpr::Literal(LiteralValue::Integer(value)) => {
1354            Some((*value).max(0).try_into().unwrap_or(usize::MAX))
1355        }
1356        _ => None,
1357    }
1358}
1359
1360fn property_access_for_var(expr: &ResolvedExpr, var: VarId) -> Option<String> {
1361    match expr {
1362        ResolvedExpr::Property { expr, property } => match &**expr {
1363            ResolvedExpr::Variable(v) if *v == var => Some(property.clone()),
1364            _ => None,
1365        },
1366        _ => None,
1367    }
1368}
1369
1370fn collect_vars_inner(expr: &ResolvedExpr, out: &mut BTreeSet<VarId>) {
1371    match expr {
1372        ResolvedExpr::Variable(v) => {
1373            out.insert(*v);
1374        }
1375        ResolvedExpr::Property { expr, .. } => collect_vars_inner(expr, out),
1376        ResolvedExpr::Binary { lhs, rhs, .. } => {
1377            collect_vars_inner(lhs, out);
1378            collect_vars_inner(rhs, out);
1379        }
1380        ResolvedExpr::Unary { expr, .. } => collect_vars_inner(expr, out),
1381        ResolvedExpr::Function { args, .. } => {
1382            for arg in args {
1383                collect_vars_inner(arg, out);
1384            }
1385        }
1386        ResolvedExpr::List(items) => {
1387            for item in items {
1388                collect_vars_inner(item, out);
1389            }
1390        }
1391        ResolvedExpr::Map(items) => {
1392            for (_, v) in items {
1393                collect_vars_inner(v, out);
1394            }
1395        }
1396        ResolvedExpr::Case {
1397            input,
1398            alternatives,
1399            else_expr,
1400        } => {
1401            if let Some(e) = input {
1402                collect_vars_inner(e, out);
1403            }
1404            for (w, t) in alternatives {
1405                collect_vars_inner(w, out);
1406                collect_vars_inner(t, out);
1407            }
1408            if let Some(e) = else_expr {
1409                collect_vars_inner(e, out);
1410            }
1411        }
1412        ResolvedExpr::ListPredicate {
1413            variable,
1414            list,
1415            predicate,
1416            ..
1417        } => {
1418            out.insert(*variable);
1419            collect_vars_inner(list, out);
1420            collect_vars_inner(predicate, out);
1421        }
1422        ResolvedExpr::ListComprehension {
1423            variable,
1424            list,
1425            filter,
1426            map_expr,
1427            ..
1428        } => {
1429            out.insert(*variable);
1430            collect_vars_inner(list, out);
1431            if let Some(f) = filter {
1432                collect_vars_inner(f, out);
1433            }
1434            if let Some(m) = map_expr {
1435                collect_vars_inner(m, out);
1436            }
1437        }
1438        ResolvedExpr::Reduce {
1439            accumulator,
1440            init,
1441            variable,
1442            list,
1443            expr,
1444            ..
1445        } => {
1446            out.insert(*accumulator);
1447            out.insert(*variable);
1448            collect_vars_inner(init, out);
1449            collect_vars_inner(list, out);
1450            collect_vars_inner(expr, out);
1451        }
1452        ResolvedExpr::Index { expr, index } => {
1453            collect_vars_inner(expr, out);
1454            collect_vars_inner(index, out);
1455        }
1456        ResolvedExpr::Slice { expr, from, to } => {
1457            collect_vars_inner(expr, out);
1458            if let Some(f) = from {
1459                collect_vars_inner(f, out);
1460            }
1461            if let Some(t) = to {
1462                collect_vars_inner(t, out);
1463            }
1464        }
1465        ResolvedExpr::MapProjection { base, selectors } => {
1466            collect_vars_inner(base, out);
1467            for sel in selectors {
1468                if let lora_analyzer::ResolvedMapSelector::Literal(_, e) = sel {
1469                    collect_vars_inner(e, out);
1470                }
1471            }
1472        }
1473        _ => {}
1474    }
1475}
1476
1477#[cfg(test)]
1478mod tests {
1479    use super::*;
1480    use lora_store::GraphStats;
1481
1482    fn stats_with_label(label: &str, total: usize, distinct: Option<usize>) -> GraphStats {
1483        let mut s = GraphStats {
1484            node_count: total,
1485            ..Default::default()
1486        };
1487        s.nodes_by_label.insert(label.to_string(), total);
1488        if let Some(d) = distinct {
1489            s.node_distinct_values
1490                .insert((label.to_string(), "id".to_string()), d);
1491        }
1492        s
1493    }
1494
1495    fn person_labels() -> Vec<Vec<String>> {
1496        vec![vec!["Person".to_string()]]
1497    }
1498
1499    fn lit_int(v: i64) -> ResolvedExpr {
1500        ResolvedExpr::Literal(LiteralValue::Integer(v))
1501    }
1502
1503    fn lit_str(s: &str) -> ResolvedExpr {
1504        ResolvedExpr::Literal(LiteralValue::String(s.to_string()))
1505    }
1506
1507    fn lit_float(v: f64) -> ResolvedExpr {
1508        ResolvedExpr::Literal(LiteralValue::Float(v))
1509    }
1510
1511    fn point_lonlat(lon: f64, lat: f64) -> ResolvedExpr {
1512        ResolvedExpr::Function {
1513            function: lora_analyzer::FunctionId::builtin("cast.to")
1514                .expect("cast.to builtin exists"),
1515            distinct: false,
1516            args: vec![
1517                ResolvedExpr::Map(vec![
1518                    ("longitude".to_string(), lit_float(lon)),
1519                    ("latitude".to_string(), lit_float(lat)),
1520                ]),
1521                ResolvedExpr::Literal(LiteralValue::TypeName("POINT".to_string())),
1522            ],
1523        }
1524    }
1525
1526    // ---------- score_logical_op ----------
1527
1528    #[test]
1529    fn score_label_scan_returns_label_count() {
1530        let stats = stats_with_label("Person", 1_000, None);
1531        let op = LogicalOp::NodeScan(NodeScan {
1532            input: None,
1533            var: VarId(0),
1534            labels: person_labels(),
1535        });
1536        assert_eq!(score_logical_op(&op, &stats), Some(1_000));
1537    }
1538
1539    #[test]
1540    fn score_property_scan_uses_distinct() {
1541        // 1000 nodes, 100 distinct values for `id`: ~10 per value.
1542        let stats = stats_with_label("Person", 1_000, Some(100));
1543        let op = LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
1544            input: None,
1545            var: VarId(0),
1546            labels: person_labels(),
1547            key: "id".to_string(),
1548            value: lit_int(7),
1549            in_list: false,
1550        });
1551        assert_eq!(score_logical_op(&op, &stats), Some(10));
1552    }
1553
1554    #[test]
1555    fn score_property_scan_high_distinct_beats_label_scan() {
1556        // distinct == total → uniform-distribution heuristic gives 1
1557        // estimated row, well below the 100-row label scan.
1558        let stats = stats_with_label("Person", 100, Some(100));
1559        let label_score = score_logical_op(
1560            &LogicalOp::NodeScan(NodeScan {
1561                input: None,
1562                var: VarId(0),
1563                labels: person_labels(),
1564            }),
1565            &stats,
1566        );
1567        let property_score = score_logical_op(
1568            &LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
1569                input: None,
1570                var: VarId(0),
1571                labels: person_labels(),
1572                key: "id".to_string(),
1573                value: lit_int(7),
1574                in_list: false,
1575            }),
1576            &stats,
1577        );
1578        assert!(property_score < label_score);
1579    }
1580
1581    #[test]
1582    fn score_returns_none_without_label_stats() {
1583        // Empty stats: every estimator should fail open so the optimizer
1584        // can fall back to the legacy "commit any matching rewrite"
1585        // behaviour through `improves_over(None, None)`.
1586        let stats = GraphStats::default();
1587        let op = LogicalOp::NodeScan(NodeScan {
1588            input: None,
1589            var: VarId(0),
1590            labels: person_labels(),
1591        });
1592        assert_eq!(score_logical_op(&op, &stats), None);
1593    }
1594
1595    // ---------- improves_over ----------
1596
1597    #[test]
1598    fn improves_over_legacy_fallback_when_both_unknown() {
1599        // `(None, None)` must resolve to "commit the rewrite" — that is
1600        // the contract for environments without stats.
1601        assert!(improves_over(None, None));
1602    }
1603
1604    #[test]
1605    fn improves_over_keeps_baseline_when_candidate_unknown() {
1606        assert!(!improves_over(None, Some(10)));
1607    }
1608
1609    #[test]
1610    fn improves_over_strictly_better_or_equal_wins() {
1611        assert!(improves_over(Some(5), Some(10)));
1612        assert!(improves_over(Some(10), Some(10)));
1613        assert!(!improves_over(Some(11), Some(10)));
1614    }
1615
1616    // ---------- pick_best_candidate ----------
1617
1618    #[test]
1619    fn pick_best_candidate_picks_lowest_score() {
1620        let stats = stats_with_label("Person", 1_200, Some(100));
1621        let original = LogicalOp::NodeScan(NodeScan {
1622            input: None,
1623            var: VarId(0),
1624            labels: person_labels(),
1625        });
1626        // property scan ~12 rows (1200 / 100), range scan ~400 rows.
1627        let candidates = vec![
1628            LogicalOp::NodeByPropertyRangeScan(NodeByPropertyRangeScan {
1629                input: None,
1630                var: VarId(0),
1631                labels: person_labels(),
1632                key: "age".to_string(),
1633                lo: Some(lit_int(30)),
1634                lo_inclusive: false,
1635                hi: None,
1636                hi_inclusive: false,
1637                order: None,
1638            }),
1639            LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
1640                input: None,
1641                var: VarId(0),
1642                labels: person_labels(),
1643                key: "id".to_string(),
1644                value: lit_int(7),
1645                in_list: false,
1646            }),
1647        ];
1648        let pick = pick_best_candidate(&original, candidates, &stats).expect("expected a pick");
1649        assert!(matches!(pick, LogicalOp::NodeByPropertyScan(_)));
1650    }
1651
1652    #[test]
1653    fn pick_best_candidate_returns_none_when_no_candidate_improves() {
1654        // Stats put baseline at 1 row (1 node, distinct=1) and the
1655        // property scan also at 1 — but no rewrite improves, and we'd
1656        // still commit the candidate because of the s<=b fallback.
1657        // To prove we *can* return None, give the candidate a higher
1658        // score than baseline.
1659        let mut stats = stats_with_label("Person", 1, None);
1660        // 1 Person, 1 distinct → property scan = 1, label scan = 1.
1661        // Add a rare label so label_estimate gives 1 but the candidate
1662        // has no distinct entry → unknown score → can't improve.
1663        stats.nodes_by_label.insert("Tiny".to_string(), 1);
1664        let original = LogicalOp::NodeScan(NodeScan {
1665            input: None,
1666            var: VarId(0),
1667            labels: vec![vec!["Tiny".to_string()]],
1668        });
1669        // Candidate references a different label → label_estimate fails
1670        // → score is None. Baseline is Some(1). improves_over(None,
1671        // Some(1)) = false → no rewrite.
1672        let candidates = vec![LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
1673            input: None,
1674            var: VarId(0),
1675            labels: vec![vec!["Missing".to_string()]],
1676            key: "id".to_string(),
1677            value: lit_int(7),
1678            in_list: false,
1679        })];
1680        assert!(pick_best_candidate(&original, candidates, &stats).is_none());
1681    }
1682
1683    // ---------- tautology guards ----------
1684
1685    #[test]
1686    fn unbounded_low_range_is_tautological() {
1687        let bounds = RangeBounds {
1688            key: "age".to_string(),
1689            lo: Some(lit_int(i64::MIN)),
1690            lo_inclusive: false,
1691            hi: None,
1692            hi_inclusive: false,
1693        };
1694        assert!(is_tautological_range(&bounds));
1695    }
1696
1697    #[test]
1698    fn unbounded_high_range_is_tautological() {
1699        let bounds = RangeBounds {
1700            key: "age".to_string(),
1701            lo: None,
1702            lo_inclusive: false,
1703            hi: Some(lit_int(i64::MAX)),
1704            hi_inclusive: false,
1705        };
1706        assert!(is_tautological_range(&bounds));
1707    }
1708
1709    #[test]
1710    fn doubly_unbounded_range_is_tautological() {
1711        let bounds = RangeBounds {
1712            key: "age".to_string(),
1713            lo: Some(lit_int(i64::MIN)),
1714            lo_inclusive: false,
1715            hi: Some(lit_int(i64::MAX)),
1716            hi_inclusive: false,
1717        };
1718        assert!(is_tautological_range(&bounds));
1719    }
1720
1721    #[test]
1722    fn ordinary_range_is_not_tautological() {
1723        let bounds = RangeBounds {
1724            key: "age".to_string(),
1725            lo: Some(lit_int(0)),
1726            lo_inclusive: false,
1727            hi: Some(lit_int(100)),
1728            hi_inclusive: false,
1729        };
1730        assert!(!is_tautological_range(&bounds));
1731    }
1732
1733    #[test]
1734    fn empty_string_starts_with_is_tautological() {
1735        let candidate = TextCandidate {
1736            key: "name".to_string(),
1737            predicate: TextPredicate::StartsWith,
1738            query: lit_str(""),
1739        };
1740        assert!(is_tautological_text(&candidate));
1741    }
1742
1743    #[test]
1744    fn nonempty_string_starts_with_is_not_tautological() {
1745        let candidate = TextCandidate {
1746            key: "name".to_string(),
1747            predicate: TextPredicate::StartsWith,
1748            query: lit_str("A"),
1749        };
1750        assert!(!is_tautological_text(&candidate));
1751    }
1752
1753    #[test]
1754    fn world_bbox_is_tautological() {
1755        let candidate = PointCandidate {
1756            key: "loc".to_string(),
1757            predicate: PointPredicate::WithinBBox {
1758                lower_left: point_lonlat(-180.0, -90.0),
1759                upper_right: point_lonlat(180.0, 90.0),
1760            },
1761        };
1762        assert!(is_tautological_point(&candidate));
1763    }
1764
1765    #[test]
1766    fn city_bbox_is_not_tautological() {
1767        let candidate = PointCandidate {
1768            key: "loc".to_string(),
1769            predicate: PointPredicate::WithinBBox {
1770                lower_left: point_lonlat(4.7, 52.3),
1771                upper_right: point_lonlat(5.0, 52.5),
1772            },
1773        };
1774        assert!(!is_tautological_point(&candidate));
1775    }
1776
1777    #[test]
1778    fn distance_bbox_is_never_tautological() {
1779        // Distance probes always narrow — no constant form folds them
1780        // into "all rows".
1781        let candidate = PointCandidate {
1782            key: "loc".to_string(),
1783            predicate: PointPredicate::WithinDistance {
1784                center: point_lonlat(0.0, 0.0),
1785                max_distance: lit_int(1_000_000_000),
1786                inclusive: true,
1787            },
1788        };
1789        assert!(!is_tautological_point(&candidate));
1790    }
1791}