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        if !is_tautological_text(&candidate)
571            && first_simple_label(&scan.labels)
572                .is_some_and(|label| stats.has_node_text_index(label, &candidate.key))
573        {
574            out.push(LogicalOp::NodeByTextScan(NodeByTextScan {
575                input: scan.input,
576                var: scan.var,
577                labels: scan.labels.clone(),
578                key: candidate.key,
579                predicate: candidate.predicate,
580                query: candidate.query,
581            }));
582        }
583    }
584
585    if let Some(candidate) = point_predicate_for_var(predicate, scan.var) {
586        if !is_tautological_point(&candidate)
587            && first_simple_label(&scan.labels)
588                .is_some_and(|label| stats.has_node_point_index(label, &candidate.key))
589        {
590            out.push(LogicalOp::NodeByPointScan(NodeByPointScan {
591                input: scan.input,
592                var: scan.var,
593                labels: scan.labels.clone(),
594                key: candidate.key,
595                predicate: candidate.predicate,
596            }));
597        }
598    }
599
600    out
601}
602
603/// Build every applicable rel-targeted index rewrite for a
604/// `Filter(Expand(NodeScan, …), pred-on-rel-var)` site. Mirrors
605/// [`collect_index_candidates`] for nodes; tautological predicates are
606/// dropped through the same `is_tautological_*` helpers.
607fn collect_rel_index_candidates(
608    expand: &Expand,
609    predicate: &ResolvedExpr,
610    rel_var: VarId,
611    input: Option<PlanNodeId>,
612    stats: &GraphStats,
613) -> Vec<LogicalOp> {
614    let mut out = Vec::new();
615
616    if let Some(candidate) = text_predicate_for_var(predicate, rel_var) {
617        if !is_tautological_text(&candidate)
618            && rel_types_have_index(&expand.types, |ty| {
619                stats.has_relationship_text_index(ty, &candidate.key)
620            })
621        {
622            out.push(LogicalOp::RelByTextScan(RelByTextScan {
623                input,
624                src: expand.src,
625                rel: rel_var,
626                dst: expand.dst,
627                types: expand.types.clone(),
628                direction: expand.direction,
629                key: candidate.key,
630                predicate: candidate.predicate,
631                query: candidate.query,
632            }));
633        }
634    }
635
636    if let Some(bounds) = collect_range_bounds(predicate, rel_var) {
637        if !is_tautological_range(&bounds)
638            && rel_types_have_index(&expand.types, |ty| {
639                stats.has_relationship_range_index(ty, &bounds.key)
640            })
641        {
642            out.push(LogicalOp::RelByPropertyRangeScan(RelByPropertyRangeScan {
643                input,
644                src: expand.src,
645                rel: rel_var,
646                dst: expand.dst,
647                types: expand.types.clone(),
648                direction: expand.direction,
649                key: bounds.key,
650                lo: bounds.lo,
651                lo_inclusive: bounds.lo_inclusive,
652                hi: bounds.hi,
653                hi_inclusive: bounds.hi_inclusive,
654            }));
655        }
656    }
657
658    if let Some(candidate) = point_predicate_for_var(predicate, rel_var) {
659        if !is_tautological_point(&candidate)
660            && rel_types_have_index(&expand.types, |ty| {
661                stats.has_relationship_point_index(ty, &candidate.key)
662            })
663        {
664            out.push(LogicalOp::RelByPointScan(RelByPointScan {
665                input,
666                src: expand.src,
667                rel: rel_var,
668                dst: expand.dst,
669                types: expand.types.clone(),
670                direction: expand.direction,
671                key: candidate.key,
672                predicate: candidate.predicate,
673            }));
674        }
675    }
676
677    out
678}
679
680fn rel_types_have_index<F>(types: &[String], mut has_index: F) -> bool
681where
682    F: FnMut(&str) -> bool,
683{
684    !types.is_empty() && types.iter().all(|ty| has_index(ty))
685}
686
687/// Choose the cheapest replacement for the original `Filter(NodeScan)`
688/// input among `candidates`, breaking ties by the caller-provided
689/// collection order so behaviour is deterministic. Returns
690/// `None` when no candidate is strictly cheaper than the original — in
691/// that case the caller leaves the plan unchanged.
692fn pick_best_candidate(
693    original: &LogicalOp,
694    candidates: Vec<LogicalOp>,
695    stats: &GraphStats,
696) -> Option<LogicalOp> {
697    if candidates.is_empty() {
698        return None;
699    }
700
701    let baseline = score_logical_op(original, stats);
702
703    let mut best: Option<(LogicalOp, Option<u64>)> = None;
704    for candidate in candidates {
705        let score = score_logical_op(&candidate, stats);
706        if !improves_over(score, baseline) {
707            continue;
708        }
709        let take = match &best {
710            None => true,
711            Some((_, current_best)) => is_cheaper(score, *current_best),
712        };
713        if take {
714            best = Some((candidate, score));
715        }
716    }
717
718    best.map(|(op, _)| op)
719}
720
721/// `true` when committing a candidate with `score` is at least as good
722/// as keeping the original (`baseline`). With unknown stats the
723/// optimizer keeps its pre-cost-model behaviour: a `None`/`None`
724/// comparison still commits already-collected candidates. Catalog
725/// checks happen before this point for range/text/point operators.
726fn improves_over(score: Option<u64>, baseline: Option<u64>) -> bool {
727    match (score, baseline) {
728        (Some(s), Some(b)) => s <= b,
729        (Some(_), None) => true,
730        (None, Some(_)) => false,
731        (None, None) => true,
732    }
733}
734
735fn is_cheaper(score: Option<u64>, current_best: Option<u64>) -> bool {
736    match (score, current_best) {
737        (Some(s), Some(b)) => s < b,
738        (Some(_), None) => true,
739        (None, _) => false,
740    }
741}
742
743/// Estimated row count produced by `op`. `None` means "no information"
744/// — typically because the relevant labels or property are not in the
745/// stats snapshot. Mirrors the per-operator estimates that `EXPLAIN`
746/// surfaces in [`crate::plan_tree::PlanTreeNode::estimated_rows`], so
747/// the optimizer's pick agrees with what users see in the plan tree.
748fn score_logical_op(op: &LogicalOp, stats: &GraphStats) -> Option<u64> {
749    match op {
750        LogicalOp::NodeScan(scan) => match label_estimate(&scan.labels, stats) {
751            Some(rows) => Some(rows),
752            None if scan.labels.is_empty() => Some(stats.node_count as u64),
753            None => None,
754        },
755        LogicalOp::NodeByPropertyScan(scan) => {
756            let label = first_simple_label(&scan.labels)?;
757            let per_value = stats.estimate_node_property_equality(label, &scan.key)?;
758            if !scan.in_list {
759                return Some(per_value);
760            }
761            // A literal list costs one seek per element; a parameter or
762            // other list-valued expression is assumed short.
763            let elements = match &scan.value {
764                ResolvedExpr::List(items) => items.len().max(1) as u64,
765                _ => 1,
766            };
767            Some(per_value.saturating_mul(elements))
768        }
769        LogicalOp::NodeByPropertyRangeScan(scan) => {
770            // Conservative one-third selectivity: a one-sided range
771            // typically narrows by less than half; a two-sided range
772            // narrows further. Better than full-label scan for any
773            // useful range, never worse than `label_count`.
774            let base = label_estimate(&scan.labels, stats)?;
775            let denom = match (scan.lo.is_some(), scan.hi.is_some()) {
776                (true, true) => 4,
777                _ => 3,
778            };
779            Some(base.div_ceil(denom))
780        }
781        LogicalOp::NodeByTextScan(scan) => {
782            let base = label_estimate(&scan.labels, stats)?;
783            // Prefix/suffix probes typically narrow more than CONTAINS.
784            let denom = match scan.predicate {
785                TextPredicate::StartsWith | TextPredicate::EndsWith => 4,
786                TextPredicate::Contains => 2,
787            };
788            Some(base.div_ceil(denom))
789        }
790        LogicalOp::NodeByPointScan(scan) => {
791            let base = label_estimate(&scan.labels, stats)?;
792            // Spatial probes: bbox/distance usually returns a small
793            // fraction of the labelled set.
794            Some(base.div_ceil(5))
795        }
796        LogicalOp::Filter(_) => None,
797        LogicalOp::Expand(expand) => {
798            // Used as the baseline when evaluating rel-index rewrites.
799            // Without per-edge histograms we approximate edges-of-type
800            // by `relationship_type_count`, falling back to the global
801            // relationship total when no type is named.
802            let count = if expand.types.is_empty() {
803                stats.relationship_count as u64
804            } else {
805                let mut total: u64 = 0;
806                for ty in &expand.types {
807                    total = total.saturating_add(stats.relationship_type_count(ty)?);
808                }
809                total
810            };
811            // Undirected expansion produces both orientations.
812            Some(match expand.direction {
813                lora_ast::Direction::Undirected => count.saturating_mul(2),
814                _ => count,
815            })
816        }
817        LogicalOp::RelByPropertyRangeScan(scan) => {
818            let base = rel_type_estimate(&scan.types, stats)?;
819            let denom = match (scan.lo.is_some(), scan.hi.is_some()) {
820                (true, true) => 4,
821                _ => 3,
822            };
823            let est = base.div_ceil(denom);
824            Some(match scan.direction {
825                lora_ast::Direction::Undirected => est.saturating_mul(2),
826                _ => est,
827            })
828        }
829        LogicalOp::RelByTextScan(scan) => {
830            let base = rel_type_estimate(&scan.types, stats)?;
831            let denom = match scan.predicate {
832                TextPredicate::StartsWith | TextPredicate::EndsWith => 4,
833                TextPredicate::Contains => 2,
834            };
835            let est = base.div_ceil(denom);
836            Some(match scan.direction {
837                lora_ast::Direction::Undirected => est.saturating_mul(2),
838                _ => est,
839            })
840        }
841        LogicalOp::RelByPointScan(scan) => {
842            let base = rel_type_estimate(&scan.types, stats)?;
843            let est = base.div_ceil(5);
844            Some(match scan.direction {
845                lora_ast::Direction::Undirected => est.saturating_mul(2),
846                _ => est,
847            })
848        }
849        _ => None,
850    }
851}
852
853fn rel_type_estimate(types: &[String], stats: &GraphStats) -> Option<u64> {
854    if types.is_empty() {
855        return Some(stats.relationship_count as u64);
856    }
857    let mut total: u64 = 0;
858    for ty in types {
859        total = total.saturating_add(stats.relationship_type_count(ty)?);
860    }
861    Some(total)
862}
863
864/// Return the count of nodes covered by a `labels` group, taking the
865/// first DNF disjunction's first literal. Mirrors `labels_estimate`
866/// from `lora-database/src/database/explain.rs` — keeping the two in
867/// sync ensures `EXPLAIN` and the optimizer agree.
868fn label_estimate(labels: &[Vec<String>], stats: &GraphStats) -> Option<u64> {
869    let label = first_simple_label(labels)?;
870    stats.label_count(label)
871}
872
873fn first_simple_label(labels: &[Vec<String>]) -> Option<&str> {
874    labels.first()?.first().map(String::as_str)
875}
876
877fn is_tautological_range(bounds: &RangeBounds) -> bool {
878    let lo_open = match (&bounds.lo, bounds.lo_inclusive) {
879        (None, _) => true,
880        (Some(expr), false) => matches!(
881            expr,
882            ResolvedExpr::Literal(LiteralValue::Integer(v)) if *v == i64::MIN
883        ),
884        (Some(_), true) => false,
885    };
886    let hi_open = match (&bounds.hi, bounds.hi_inclusive) {
887        (None, _) => true,
888        (Some(expr), false) => matches!(
889            expr,
890            ResolvedExpr::Literal(LiteralValue::Integer(v)) if *v == i64::MAX
891        ),
892        (Some(_), true) => false,
893    };
894    lo_open && hi_open
895}
896
897fn is_tautological_text(candidate: &TextCandidate) -> bool {
898    matches!(
899        &candidate.query,
900        ResolvedExpr::Literal(LiteralValue::String(s)) if s.is_empty()
901    )
902}
903
904fn is_tautological_point(candidate: &PointCandidate) -> bool {
905    match &candidate.predicate {
906        PointPredicate::WithinBBox {
907            lower_left,
908            upper_right,
909        } => is_world_bbox(lower_left, upper_right),
910        PointPredicate::WithinDistance { .. } => false,
911    }
912}
913
914/// `{longitude: -180, latitude: -90}::POINT` to
915/// `{longitude: 180, latitude: 90}::POINT` covers every WGS-84 point.
916/// Detecting the literal form avoids a trigram lookup that would yield
917/// every indexed row only to be re-filtered to the same set.
918fn is_world_bbox(lower_left: &ResolvedExpr, upper_right: &ResolvedExpr) -> bool {
919    /// Cypher parses `-180` as `Unary{Neg, Literal(180)}`, not as a
920    /// negative integer literal — peel one such layer so the
921    /// world-bbox detection works on natural query forms.
922    fn const_number(expr: &ResolvedExpr) -> Option<f64> {
923        match expr {
924            ResolvedExpr::Literal(LiteralValue::Float(v)) => Some(*v),
925            ResolvedExpr::Literal(LiteralValue::Integer(v)) => Some(*v as f64),
926            ResolvedExpr::Unary {
927                op: lora_ast::UnaryOp::Neg,
928                expr,
929            } => const_number(expr).map(|v| -v),
930            ResolvedExpr::Unary {
931                op: lora_ast::UnaryOp::Pos,
932                expr,
933            } => const_number(expr),
934            _ => None,
935        }
936    }
937
938    fn point_lon_lat(expr: &ResolvedExpr) -> Option<(f64, f64)> {
939        let items = point_literal_map(expr)?;
940        let mut lon: Option<f64> = None;
941        let mut lat: Option<f64> = None;
942        for (key, value) in items {
943            let n = const_number(value)?;
944            match key.as_str() {
945                "longitude" | "x" => lon = Some(n),
946                "latitude" | "y" => lat = Some(n),
947                _ => {}
948            }
949        }
950        Some((lon?, lat?))
951    }
952
953    let Some((ll_lon, ll_lat)) = point_lon_lat(lower_left) else {
954        return false;
955    };
956    let Some((ur_lon, ur_lat)) = point_lon_lat(upper_right) else {
957        return false;
958    };
959    ll_lon <= -180.0 && ll_lat <= -90.0 && ur_lon >= 180.0 && ur_lat >= 90.0
960}
961
962fn point_literal_map(expr: &ResolvedExpr) -> Option<&Vec<(String, ResolvedExpr)>> {
963    let ResolvedExpr::Function { function, args, .. } = expr else {
964        return None;
965    };
966    if function.eq_ignore_ascii_case("geo.point") && args.len() == 1 {
967        let ResolvedExpr::Map(items) = &args[0] else {
968            return None;
969        };
970        return Some(items);
971    }
972    if function.eq_ignore_ascii_case("cast.to") && args.len() == 2 {
973        let ResolvedExpr::Map(items) = &args[0] else {
974            return None;
975        };
976        let ResolvedExpr::Literal(LiteralValue::TypeName(target)) = &args[1] else {
977            return None;
978        };
979        if target.eq_ignore_ascii_case("POINT") {
980            return Some(items);
981        }
982    }
983    None
984}
985
986struct RangeBounds {
987    key: String,
988    lo: Option<ResolvedExpr>,
989    lo_inclusive: bool,
990    hi: Option<ResolvedExpr>,
991    hi_inclusive: bool,
992}
993
994/// Walk an AND-tree and collect any `var.prop CMP literal` bounds.
995/// Returns `None` if no comparison touches `var.prop` for a single
996/// property key — we don't try to combine multi-property bounds in v1.
997fn collect_range_bounds(predicate: &ResolvedExpr, var: VarId) -> Option<RangeBounds> {
998    let mut key: Option<String> = None;
999    let mut lo: Option<ResolvedExpr> = None;
1000    let mut lo_inclusive = false;
1001    let mut hi: Option<ResolvedExpr> = None;
1002    let mut hi_inclusive = false;
1003    let mut any = false;
1004
1005    walk_and_for_range(
1006        predicate,
1007        var,
1008        &mut |found_key, side, value, inclusive| match side {
1009            RangeSide::Lower => {
1010                if key.as_deref().map(|k| k != found_key).unwrap_or(false) {
1011                    return;
1012                }
1013                key = Some(found_key.to_string());
1014                if lo
1015                    .as_ref()
1016                    .map(|current| lower_bound_is_tighter(&value, inclusive, current, lo_inclusive))
1017                    .unwrap_or(true)
1018                {
1019                    lo = Some(value);
1020                    lo_inclusive = inclusive;
1021                }
1022                any = true;
1023            }
1024            RangeSide::Upper => {
1025                if key.as_deref().map(|k| k != found_key).unwrap_or(false) {
1026                    return;
1027                }
1028                key = Some(found_key.to_string());
1029                if hi
1030                    .as_ref()
1031                    .map(|current| upper_bound_is_tighter(&value, inclusive, current, hi_inclusive))
1032                    .unwrap_or(true)
1033                {
1034                    hi = Some(value);
1035                    hi_inclusive = inclusive;
1036                }
1037                any = true;
1038            }
1039        },
1040    );
1041
1042    if !any {
1043        return None;
1044    }
1045    Some(RangeBounds {
1046        key: key?,
1047        lo,
1048        lo_inclusive,
1049        hi,
1050        hi_inclusive,
1051    })
1052}
1053
1054#[derive(Clone, Copy)]
1055enum RangeSide {
1056    Lower,
1057    Upper,
1058}
1059
1060fn lower_bound_is_tighter(
1061    candidate: &ResolvedExpr,
1062    candidate_inclusive: bool,
1063    current: &ResolvedExpr,
1064    current_inclusive: bool,
1065) -> bool {
1066    match compare_literal_bounds(candidate, current) {
1067        Some(std::cmp::Ordering::Greater) => true,
1068        Some(std::cmp::Ordering::Equal) => !candidate_inclusive && current_inclusive,
1069        _ => false,
1070    }
1071}
1072
1073fn upper_bound_is_tighter(
1074    candidate: &ResolvedExpr,
1075    candidate_inclusive: bool,
1076    current: &ResolvedExpr,
1077    current_inclusive: bool,
1078) -> bool {
1079    match compare_literal_bounds(candidate, current) {
1080        Some(std::cmp::Ordering::Less) => true,
1081        Some(std::cmp::Ordering::Equal) => !candidate_inclusive && current_inclusive,
1082        _ => false,
1083    }
1084}
1085
1086fn compare_literal_bounds(lhs: &ResolvedExpr, rhs: &ResolvedExpr) -> Option<std::cmp::Ordering> {
1087    match (literal_number(lhs), literal_number(rhs)) {
1088        (Some(a), Some(b)) => return a.partial_cmp(&b),
1089        (Some(_), None) | (None, Some(_)) => return None,
1090        (None, None) => {}
1091    }
1092
1093    match (lhs, rhs) {
1094        (
1095            ResolvedExpr::Literal(LiteralValue::String(a)),
1096            ResolvedExpr::Literal(LiteralValue::String(b)),
1097        ) => Some(a.cmp(b)),
1098        _ => None,
1099    }
1100}
1101
1102fn literal_number(expr: &ResolvedExpr) -> Option<f64> {
1103    match expr {
1104        ResolvedExpr::Literal(LiteralValue::Integer(v)) => Some(*v as f64),
1105        ResolvedExpr::Literal(LiteralValue::Float(v)) => Some(*v),
1106        ResolvedExpr::Unary {
1107            op: lora_ast::UnaryOp::Neg,
1108            expr,
1109        } => literal_number(expr).map(|v| -v),
1110        ResolvedExpr::Unary {
1111            op: lora_ast::UnaryOp::Pos,
1112            expr,
1113        } => literal_number(expr),
1114        _ => None,
1115    }
1116}
1117
1118fn walk_and_for_range<F>(predicate: &ResolvedExpr, var: VarId, visit: &mut F)
1119where
1120    F: FnMut(&str, RangeSide, ResolvedExpr, bool),
1121{
1122    if let ResolvedExpr::Binary {
1123        lhs,
1124        op: BinaryOp::And,
1125        rhs,
1126    } = predicate
1127    {
1128        walk_and_for_range(lhs, var, visit);
1129        walk_and_for_range(rhs, var, visit);
1130        return;
1131    }
1132
1133    let ResolvedExpr::Binary { lhs, op, rhs } = predicate else {
1134        return;
1135    };
1136
1137    let (side, inclusive) = match op {
1138        BinaryOp::Gt => (RangeSide::Lower, false),
1139        BinaryOp::Ge => (RangeSide::Lower, true),
1140        BinaryOp::Lt => (RangeSide::Upper, false),
1141        BinaryOp::Le => (RangeSide::Upper, true),
1142        _ => return,
1143    };
1144
1145    if let Some(key) = property_access_for_var(lhs, var) {
1146        if !collect_vars(rhs).contains(&var) {
1147            visit(&key, side, (**rhs).clone(), inclusive);
1148            return;
1149        }
1150    }
1151    if let Some(key) = property_access_for_var(rhs, var) {
1152        if !collect_vars(lhs).contains(&var) {
1153            // Mirror `value CMP var.prop` to `var.prop FLIPPED_CMP value`.
1154            let flipped = match side {
1155                RangeSide::Lower => RangeSide::Upper,
1156                RangeSide::Upper => RangeSide::Lower,
1157            };
1158            visit(&key, flipped, (**lhs).clone(), inclusive);
1159        }
1160    }
1161}
1162
1163struct TextCandidate {
1164    key: String,
1165    predicate: TextPredicate,
1166    query: ResolvedExpr,
1167}
1168
1169struct PointCandidate {
1170    key: String,
1171    predicate: PointPredicate,
1172}
1173
1174fn point_predicate_for_var(predicate: &ResolvedExpr, var: VarId) -> Option<PointCandidate> {
1175    if let ResolvedExpr::Binary {
1176        lhs,
1177        op: BinaryOp::And,
1178        rhs,
1179    } = predicate
1180    {
1181        return point_predicate_for_var(lhs, var).or_else(|| point_predicate_for_var(rhs, var));
1182    }
1183
1184    // geo.within_bbox(n.prop, ll, ur)
1185    if let ResolvedExpr::Function { function, args, .. } = predicate {
1186        if function.eq_ignore_ascii_case("geo.within_bbox") && args.len() == 3 {
1187            if let Some(key) = property_access_for_var(&args[0], var) {
1188                if !collect_vars(&args[1]).contains(&var) && !collect_vars(&args[2]).contains(&var)
1189                {
1190                    return Some(PointCandidate {
1191                        key,
1192                        predicate: PointPredicate::WithinBBox {
1193                            lower_left: args[1].clone(),
1194                            upper_right: args[2].clone(),
1195                        },
1196                    });
1197                }
1198            }
1199        }
1200    }
1201
1202    // point.distance(n.prop, c) OP d  (where OP is <, <=)
1203    let ResolvedExpr::Binary { lhs, op, rhs } = predicate else {
1204        return None;
1205    };
1206    let inclusive = match op {
1207        BinaryOp::Le => true,
1208        BinaryOp::Lt => false,
1209        // Symmetric form: d >= point.distance(n.prop, c)
1210        BinaryOp::Ge => true,
1211        BinaryOp::Gt => false,
1212        _ => return None,
1213    };
1214
1215    let (call_side, scalar_side) = match op {
1216        BinaryOp::Le | BinaryOp::Lt => ((**lhs).clone(), (**rhs).clone()),
1217        BinaryOp::Ge | BinaryOp::Gt => ((**rhs).clone(), (**lhs).clone()),
1218        _ => return None,
1219    };
1220
1221    let ResolvedExpr::Function { function, args, .. } = &call_side else {
1222        return None;
1223    };
1224    if !function.eq_ignore_ascii_case("geo.distance") {
1225        return None;
1226    }
1227    if args.len() != 2 {
1228        return None;
1229    }
1230    let key = property_access_for_var(&args[0], var)?;
1231    if collect_vars(&args[1]).contains(&var) || collect_vars(&scalar_side).contains(&var) {
1232        return None;
1233    }
1234    Some(PointCandidate {
1235        key,
1236        predicate: PointPredicate::WithinDistance {
1237            center: args[1].clone(),
1238            max_distance: scalar_side,
1239            inclusive,
1240        },
1241    })
1242}
1243
1244fn text_predicate_for_var(predicate: &ResolvedExpr, var: VarId) -> Option<TextCandidate> {
1245    let ResolvedExpr::Binary { lhs, op, rhs } = predicate else {
1246        return None;
1247    };
1248
1249    if matches!(op, BinaryOp::And) {
1250        return text_predicate_for_var(lhs, var).or_else(|| text_predicate_for_var(rhs, var));
1251    }
1252
1253    let kind = match op {
1254        BinaryOp::StartsWith => TextPredicate::StartsWith,
1255        BinaryOp::EndsWith => TextPredicate::EndsWith,
1256        BinaryOp::Contains => TextPredicate::Contains,
1257        _ => return None,
1258    };
1259
1260    let key = property_access_for_var(lhs, var)?;
1261    if collect_vars(rhs).contains(&var) {
1262        return None;
1263    }
1264    Some(TextCandidate {
1265        key,
1266        predicate: kind,
1267        query: (**rhs).clone(),
1268    })
1269}
1270
1271/// Every `var.key = value` conjunct of an AND-tree, in order, where
1272/// `value` does not mention `var`.
1273fn property_equalities_for_var(
1274    predicate: &ResolvedExpr,
1275    var: VarId,
1276) -> Vec<(String, ResolvedExpr)> {
1277    let mut out = Vec::new();
1278    for conjunct in and_conjuncts(predicate) {
1279        let ResolvedExpr::Binary {
1280            lhs,
1281            op: BinaryOp::Eq,
1282            rhs,
1283        } = conjunct
1284        else {
1285            continue;
1286        };
1287        if let Some(key) =
1288            property_access_for_var(lhs, var).filter(|_| !collect_vars(rhs).contains(&var))
1289        {
1290            out.push((key, (**rhs).clone()));
1291        } else if let Some(key) =
1292            property_access_for_var(rhs, var).filter(|_| !collect_vars(lhs).contains(&var))
1293        {
1294            out.push((key, (**lhs).clone()));
1295        }
1296    }
1297    out
1298}
1299
1300/// Every `var.key IN list` conjunct of an AND-tree, in order, where `list`
1301/// does not mention `var`.
1302fn property_in_lists_for_var(predicate: &ResolvedExpr, var: VarId) -> 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::In,
1308            rhs,
1309        } = conjunct
1310        else {
1311            continue;
1312        };
1313        if collect_vars(rhs).contains(&var) {
1314            continue;
1315        }
1316        if let Some(key) = property_access_for_var(lhs, var) {
1317            out.push((key, (**rhs).clone()));
1318        }
1319    }
1320    out
1321}
1322
1323/// Flatten an AND-tree into its conjuncts, left to right.
1324fn and_conjuncts(predicate: &ResolvedExpr) -> Vec<&ResolvedExpr> {
1325    let mut out = Vec::new();
1326    let mut stack = vec![predicate];
1327    while let Some(expr) = stack.pop() {
1328        match expr {
1329            ResolvedExpr::Binary {
1330                lhs,
1331                op: BinaryOp::And,
1332                rhs,
1333            } => {
1334                stack.push(rhs);
1335                stack.push(lhs);
1336            }
1337            other => out.push(other),
1338        }
1339    }
1340    out
1341}
1342
1343fn static_limit_bound(limit: &Limit) -> Option<usize> {
1344    let limit_rows = static_non_negative_usize(limit.limit.as_ref()?)?;
1345    let skip_rows = limit
1346        .skip
1347        .as_ref()
1348        .and_then(static_non_negative_usize)
1349        .unwrap_or(0);
1350    Some(skip_rows.saturating_add(limit_rows))
1351}
1352
1353fn limit_sort_bound(op: &LogicalOp) -> Option<(PlanNodeId, usize)> {
1354    let LogicalOp::Limit(limit) = op else {
1355        return None;
1356    };
1357
1358    static_limit_bound(limit).map(|bound| (limit.input, bound))
1359}
1360
1361fn sort_op_mut(op: &mut LogicalOp) -> Option<&mut Sort> {
1362    match op {
1363        LogicalOp::Sort(sort) => Some(sort),
1364        _ => None,
1365    }
1366}
1367
1368fn merge_top_k_bound(current: Option<usize>, bound: usize) -> Option<usize> {
1369    Some(current.map(|current| current.min(bound)).unwrap_or(bound))
1370}
1371
1372fn static_non_negative_usize(expr: &ResolvedExpr) -> Option<usize> {
1373    match expr {
1374        ResolvedExpr::Literal(LiteralValue::Integer(value)) => {
1375            Some((*value).max(0).try_into().unwrap_or(usize::MAX))
1376        }
1377        _ => None,
1378    }
1379}
1380
1381fn property_access_for_var(expr: &ResolvedExpr, var: VarId) -> Option<String> {
1382    match expr {
1383        ResolvedExpr::Property { expr, property } => match &**expr {
1384            ResolvedExpr::Variable(v) if *v == var => Some(property.clone()),
1385            _ => None,
1386        },
1387        _ => None,
1388    }
1389}
1390
1391#[cfg(test)]
1392mod tests {
1393    use super::*;
1394    use lora_store::GraphStats;
1395
1396    fn stats_with_label(label: &str, total: usize, distinct: Option<usize>) -> GraphStats {
1397        let mut s = GraphStats {
1398            node_count: total,
1399            ..Default::default()
1400        };
1401        s.nodes_by_label.insert(label.to_string(), total);
1402        if let Some(d) = distinct {
1403            s.node_distinct_values
1404                .insert((label.to_string(), "id".to_string()), d);
1405        }
1406        s
1407    }
1408
1409    fn person_labels() -> Vec<Vec<String>> {
1410        vec![vec!["Person".to_string()]]
1411    }
1412
1413    fn lit_int(v: i64) -> ResolvedExpr {
1414        ResolvedExpr::Literal(LiteralValue::Integer(v))
1415    }
1416
1417    fn lit_str(s: &str) -> ResolvedExpr {
1418        ResolvedExpr::Literal(LiteralValue::String(s.to_string()))
1419    }
1420
1421    fn lit_float(v: f64) -> ResolvedExpr {
1422        ResolvedExpr::Literal(LiteralValue::Float(v))
1423    }
1424
1425    fn point_lonlat(lon: f64, lat: f64) -> ResolvedExpr {
1426        ResolvedExpr::Function {
1427            function: lora_analyzer::FunctionId::builtin("cast.to")
1428                .expect("cast.to builtin exists"),
1429            distinct: false,
1430            args: vec![
1431                ResolvedExpr::Map(vec![
1432                    ("longitude".to_string(), lit_float(lon)),
1433                    ("latitude".to_string(), lit_float(lat)),
1434                ]),
1435                ResolvedExpr::Literal(LiteralValue::TypeName("POINT".to_string())),
1436            ],
1437        }
1438    }
1439
1440    // ---------- score_logical_op ----------
1441
1442    #[test]
1443    fn score_label_scan_returns_label_count() {
1444        let stats = stats_with_label("Person", 1_000, None);
1445        let op = LogicalOp::NodeScan(NodeScan {
1446            input: None,
1447            var: VarId(0),
1448            labels: person_labels(),
1449        });
1450        assert_eq!(score_logical_op(&op, &stats), Some(1_000));
1451    }
1452
1453    #[test]
1454    fn score_property_scan_uses_distinct() {
1455        // 1000 nodes, 100 distinct values for `id`: ~10 per value.
1456        let stats = stats_with_label("Person", 1_000, Some(100));
1457        let op = LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
1458            input: None,
1459            var: VarId(0),
1460            labels: person_labels(),
1461            key: "id".to_string(),
1462            value: lit_int(7),
1463            in_list: false,
1464        });
1465        assert_eq!(score_logical_op(&op, &stats), Some(10));
1466    }
1467
1468    #[test]
1469    fn score_property_scan_high_distinct_beats_label_scan() {
1470        // distinct == total → uniform-distribution heuristic gives 1
1471        // estimated row, well below the 100-row label scan.
1472        let stats = stats_with_label("Person", 100, Some(100));
1473        let label_score = score_logical_op(
1474            &LogicalOp::NodeScan(NodeScan {
1475                input: None,
1476                var: VarId(0),
1477                labels: person_labels(),
1478            }),
1479            &stats,
1480        );
1481        let property_score = score_logical_op(
1482            &LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
1483                input: None,
1484                var: VarId(0),
1485                labels: person_labels(),
1486                key: "id".to_string(),
1487                value: lit_int(7),
1488                in_list: false,
1489            }),
1490            &stats,
1491        );
1492        assert!(property_score < label_score);
1493    }
1494
1495    #[test]
1496    fn score_returns_none_without_label_stats() {
1497        // Empty stats: every estimator should fail open so the optimizer
1498        // can fall back to the legacy "commit any matching rewrite"
1499        // behaviour through `improves_over(None, None)`.
1500        let stats = GraphStats::default();
1501        let op = LogicalOp::NodeScan(NodeScan {
1502            input: None,
1503            var: VarId(0),
1504            labels: person_labels(),
1505        });
1506        assert_eq!(score_logical_op(&op, &stats), None);
1507    }
1508
1509    // ---------- improves_over ----------
1510
1511    #[test]
1512    fn improves_over_legacy_fallback_when_both_unknown() {
1513        // `(None, None)` must resolve to "commit the rewrite" — that is
1514        // the contract for environments without stats.
1515        assert!(improves_over(None, None));
1516    }
1517
1518    #[test]
1519    fn improves_over_keeps_baseline_when_candidate_unknown() {
1520        assert!(!improves_over(None, Some(10)));
1521    }
1522
1523    #[test]
1524    fn improves_over_strictly_better_or_equal_wins() {
1525        assert!(improves_over(Some(5), Some(10)));
1526        assert!(improves_over(Some(10), Some(10)));
1527        assert!(!improves_over(Some(11), Some(10)));
1528    }
1529
1530    // ---------- pick_best_candidate ----------
1531
1532    #[test]
1533    fn pick_best_candidate_picks_lowest_score() {
1534        let stats = stats_with_label("Person", 1_200, Some(100));
1535        let original = LogicalOp::NodeScan(NodeScan {
1536            input: None,
1537            var: VarId(0),
1538            labels: person_labels(),
1539        });
1540        // property scan ~12 rows (1200 / 100), range scan ~400 rows.
1541        let candidates = vec![
1542            LogicalOp::NodeByPropertyRangeScan(NodeByPropertyRangeScan {
1543                input: None,
1544                var: VarId(0),
1545                labels: person_labels(),
1546                key: "age".to_string(),
1547                lo: Some(lit_int(30)),
1548                lo_inclusive: false,
1549                hi: None,
1550                hi_inclusive: false,
1551                order: None,
1552            }),
1553            LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
1554                input: None,
1555                var: VarId(0),
1556                labels: person_labels(),
1557                key: "id".to_string(),
1558                value: lit_int(7),
1559                in_list: false,
1560            }),
1561        ];
1562        let pick = pick_best_candidate(&original, candidates, &stats).expect("expected a pick");
1563        assert!(matches!(pick, LogicalOp::NodeByPropertyScan(_)));
1564    }
1565
1566    #[test]
1567    fn pick_best_candidate_returns_none_when_no_candidate_improves() {
1568        // Stats put baseline at 1 row (1 node, distinct=1) and the
1569        // property scan also at 1 — but no rewrite improves, and we'd
1570        // still commit the candidate because of the s<=b fallback.
1571        // To prove we *can* return None, give the candidate a higher
1572        // score than baseline.
1573        let mut stats = stats_with_label("Person", 1, None);
1574        // 1 Person, 1 distinct → property scan = 1, label scan = 1.
1575        // Add a rare label so label_estimate gives 1 but the candidate
1576        // has no distinct entry → unknown score → can't improve.
1577        stats.nodes_by_label.insert("Tiny".to_string(), 1);
1578        let original = LogicalOp::NodeScan(NodeScan {
1579            input: None,
1580            var: VarId(0),
1581            labels: vec![vec!["Tiny".to_string()]],
1582        });
1583        // Candidate references a different label → label_estimate fails
1584        // → score is None. Baseline is Some(1). improves_over(None,
1585        // Some(1)) = false → no rewrite.
1586        let candidates = vec![LogicalOp::NodeByPropertyScan(NodeByPropertyScan {
1587            input: None,
1588            var: VarId(0),
1589            labels: vec![vec!["Missing".to_string()]],
1590            key: "id".to_string(),
1591            value: lit_int(7),
1592            in_list: false,
1593        })];
1594        assert!(pick_best_candidate(&original, candidates, &stats).is_none());
1595    }
1596
1597    // ---------- tautology guards ----------
1598
1599    #[test]
1600    fn unbounded_low_range_is_tautological() {
1601        let bounds = RangeBounds {
1602            key: "age".to_string(),
1603            lo: Some(lit_int(i64::MIN)),
1604            lo_inclusive: false,
1605            hi: None,
1606            hi_inclusive: false,
1607        };
1608        assert!(is_tautological_range(&bounds));
1609    }
1610
1611    #[test]
1612    fn unbounded_high_range_is_tautological() {
1613        let bounds = RangeBounds {
1614            key: "age".to_string(),
1615            lo: None,
1616            lo_inclusive: false,
1617            hi: Some(lit_int(i64::MAX)),
1618            hi_inclusive: false,
1619        };
1620        assert!(is_tautological_range(&bounds));
1621    }
1622
1623    #[test]
1624    fn doubly_unbounded_range_is_tautological() {
1625        let bounds = RangeBounds {
1626            key: "age".to_string(),
1627            lo: Some(lit_int(i64::MIN)),
1628            lo_inclusive: false,
1629            hi: Some(lit_int(i64::MAX)),
1630            hi_inclusive: false,
1631        };
1632        assert!(is_tautological_range(&bounds));
1633    }
1634
1635    #[test]
1636    fn ordinary_range_is_not_tautological() {
1637        let bounds = RangeBounds {
1638            key: "age".to_string(),
1639            lo: Some(lit_int(0)),
1640            lo_inclusive: false,
1641            hi: Some(lit_int(100)),
1642            hi_inclusive: false,
1643        };
1644        assert!(!is_tautological_range(&bounds));
1645    }
1646
1647    #[test]
1648    fn empty_string_starts_with_is_tautological() {
1649        let candidate = TextCandidate {
1650            key: "name".to_string(),
1651            predicate: TextPredicate::StartsWith,
1652            query: lit_str(""),
1653        };
1654        assert!(is_tautological_text(&candidate));
1655    }
1656
1657    #[test]
1658    fn nonempty_string_starts_with_is_not_tautological() {
1659        let candidate = TextCandidate {
1660            key: "name".to_string(),
1661            predicate: TextPredicate::StartsWith,
1662            query: lit_str("A"),
1663        };
1664        assert!(!is_tautological_text(&candidate));
1665    }
1666
1667    #[test]
1668    fn world_bbox_is_tautological() {
1669        let candidate = PointCandidate {
1670            key: "loc".to_string(),
1671            predicate: PointPredicate::WithinBBox {
1672                lower_left: point_lonlat(-180.0, -90.0),
1673                upper_right: point_lonlat(180.0, 90.0),
1674            },
1675        };
1676        assert!(is_tautological_point(&candidate));
1677    }
1678
1679    #[test]
1680    fn city_bbox_is_not_tautological() {
1681        let candidate = PointCandidate {
1682            key: "loc".to_string(),
1683            predicate: PointPredicate::WithinBBox {
1684                lower_left: point_lonlat(4.7, 52.3),
1685                upper_right: point_lonlat(5.0, 52.5),
1686            },
1687        };
1688        assert!(!is_tautological_point(&candidate));
1689    }
1690
1691    #[test]
1692    fn distance_bbox_is_never_tautological() {
1693        // Distance probes always narrow — no constant form folds them
1694        // into "all rows".
1695        let candidate = PointCandidate {
1696            key: "loc".to_string(),
1697            predicate: PointPredicate::WithinDistance {
1698                center: point_lonlat(0.0, 0.0),
1699                max_distance: lit_int(1_000_000_000),
1700                inclusive: true,
1701            },
1702        };
1703        assert!(!is_tautological_point(&candidate));
1704    }
1705}