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