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