use std::collections::HashMap;
use crate::query::ast::{Attr, Query, SelectItem};
use crate::query::carry::CarryLayout;
use crate::query::plan::DeferredProj;
use crate::query::plan::{PredCost, QueryPlan, StageOp};
#[derive(Debug, Default, Clone)]
pub struct SchemaStats {
#[allow(dead_code)]
pub instance_counts: HashMap<String, u64>,
}
impl SchemaStats {
#[allow(dead_code)]
pub fn count_of(&self, class: &str) -> u64 {
self.instance_counts.get(class).copied().unwrap_or(0)
}
}
pub fn pred_cost_rank(c: PredCost) -> u8 {
match c {
PredCost::Type => 0,
PredCost::Scalar => 1,
PredCost::Str => 2,
PredCost::Ref => 3,
}
}
pub fn reorder_predicates(plan: &mut QueryPlan) {
plan.where_terms.sort_by_key(|c| pred_cost_rank(c.cost));
}
pub fn pushdown_limit(plan: &mut QueryPlan) {
let safe = plan.limit.is_some()
&& !plan.order_sensitive
&& plan.late_ops.is_empty()
&& plan.union_branches.is_empty()
&& plan.from_subplan.is_none()
&& plan.in_subplans.is_empty();
plan.scan_limit = if safe { plan.limit } else { None };
}
fn select_item_is_deferrable(item: &SelectItem) -> bool {
match item {
SelectItem::Attr(Attr::RefPath { .. }) | SelectItem::Attr(Attr::RetainedHeapSize) => true,
SelectItem::Aggregate { arg, .. } => select_item_is_deferrable(arg),
_ => false,
}
}
pub fn defer_projections(plan: &mut QueryPlan, query: &Query) {
plan.deferred_projections.clear();
for (i, item) in query.select.iter().enumerate() {
if select_item_is_deferrable(item) {
plan.deferred_projections
.push(DeferredProj { select_index: i });
}
}
}
pub fn eliminate_dead_needs(plan: &mut QueryPlan) {
let mut retained = false;
let mut dominator_children = false;
let mut ref_walk = false;
let mut array_index = false;
for op in &plan.late_ops {
match op {
StageOp::JoinRetained => retained = true,
StageOp::RetainedSet { .. } => {
retained = true;
dominator_children = true;
}
StageOp::DominatorChildren { .. } | StageOp::DominatorOf => {
dominator_children = true;
}
StageOp::RefWalkResolve { .. } => ref_walk = true,
StageOp::ResolveArrayIndex => array_index = true,
StageOp::EdgeLookup { .. } | StageOp::BoundedPath { .. } => {}
StageOp::ResolveStringValues => {}
}
}
plan.needs.retained &= retained;
plan.needs.dominator_children &= dominator_children;
plan.needs.ref_walk &= ref_walk;
plan.needs.array_index &= array_index;
}
pub fn narrow_carry(plan: &mut QueryPlan) {
if let CarryLayout::IndexPlusScalars { widths } = &plan.carry {
if widths.is_empty() {
plan.carry = CarryLayout::IndexOnly;
}
}
}
pub fn order_by_selectivity(_plan: &mut QueryPlan, _stats: &SchemaStats) {
}
pub fn optimize(mut plan: QueryPlan, query: &Query, stats: &SchemaStats) -> QueryPlan {
reorder_predicates(&mut plan);
order_by_selectivity(&mut plan, stats);
pushdown_limit(&mut plan);
defer_projections(&mut plan, query);
eliminate_dead_needs(&mut plan);
narrow_carry(&mut plan);
if let Some(sub) = plan.from_subplan.take() {
plan.from_subplan = Some(Box::new(match query.from.as_subquery() {
Some(sub_ast) => optimize(*sub, sub_ast, stats),
None => *sub,
}));
}
for isp in &mut plan.in_subplans {
let inner = isp.inner.clone();
let sub = std::mem::take(&mut isp.plan);
isp.plan = optimize(sub, &inner, stats);
}
let branch_asts = &query.union_branches;
plan.union_branches = plan
.union_branches
.into_iter()
.enumerate()
.map(|(i, b)| match branch_asts.get(i) {
Some(bast) => optimize(b, bast, stats),
None => b,
})
.collect();
plan
}
#[cfg(test)]
mod tests {
use super::*;
use crate::query::ast::{Attr, CompareOp, Expr, Predicate, Value};
use crate::query::carry::CarryLayout;
use crate::query::parse::parse;
use crate::query::plan::StageOp;
use crate::query::plan::plan_query;
use crate::query::plan::{Conjunct, Phase, PredCost, StageKind};
fn pq(q: &crate::query::ast::Query) -> QueryPlan {
plan_query(q, crate::query::DEFAULT_PATH_DEPTH_CAP).unwrap()
}
fn scalar_conjunct(field: &str, cost: PredCost) -> Conjunct {
Conjunct {
pred: Predicate::Compare {
lhs: Expr::Attr(Attr::Field(field.to_string())),
op: CompareOp::Gt,
rhs: Expr::Lit(Value::Int(0)),
},
cost,
}
}
#[test]
fn reorder_sorts_cheap_first() {
let mut plan = pq(&parse("SELECT @objectId FROM java.lang.String").unwrap());
plan.where_terms = vec![
scalar_conjunct("d", PredCost::Ref),
scalar_conjunct("c", PredCost::Str),
scalar_conjunct("b", PredCost::Scalar),
scalar_conjunct("a", PredCost::Type),
];
reorder_predicates(&mut plan);
let ranks: Vec<u8> = plan
.where_terms
.iter()
.map(|c| pred_cost_rank(c.cost))
.collect();
assert!(
ranks.windows(2).all(|w| w[0] <= w[1]),
"expected non-decreasing ranks, got: {:?}",
ranks
);
assert_eq!(
ranks.first().copied(),
Some(0),
"first conjunct must be cheapest"
);
assert_eq!(
ranks.last().copied(),
Some(3),
"last conjunct must be most expensive"
);
}
#[test]
fn reorder_is_stable_within_cost_class() {
let mut plan = pq(&parse("SELECT @objectId FROM java.lang.String").unwrap());
plan.where_terms = vec![
scalar_conjunct("first", PredCost::Scalar),
scalar_conjunct("second", PredCost::Scalar),
];
reorder_predicates(&mut plan);
assert_eq!(plan.where_terms.len(), 2);
match &plan.where_terms[0].pred {
Predicate::Compare {
lhs: Expr::Attr(Attr::Field(name)),
..
} => {
assert_eq!(
name, "first",
"stable sort must preserve first conjunct's position"
);
}
other => panic!("unexpected predicate: {:?}", other),
}
match &plan.where_terms[1].pred {
Predicate::Compare {
lhs: Expr::Attr(Attr::Field(name)),
..
} => {
assert_eq!(
name, "second",
"stable sort must preserve second conjunct's position"
);
}
other => panic!("unexpected predicate: {:?}", other),
}
}
#[test]
fn reorder_is_idempotent() {
let mut plan = pq(&parse("SELECT @objectId FROM java.lang.String").unwrap());
plan.where_terms = vec![
scalar_conjunct("d", PredCost::Ref),
scalar_conjunct("c", PredCost::Str),
scalar_conjunct("b", PredCost::Scalar),
scalar_conjunct("a", PredCost::Type),
];
reorder_predicates(&mut plan);
let after_first: Vec<PredCost> = plan.where_terms.iter().map(|c| c.cost).collect();
reorder_predicates(&mut plan);
let after_second: Vec<PredCost> = plan.where_terms.iter().map(|c| c.cost).collect();
assert_eq!(
after_first, after_second,
"reorder_predicates must be idempotent"
);
}
#[test]
fn reorder_empty_where_is_noop() {
let mut plan = pq(&parse("SELECT * FROM java.lang.String").unwrap());
assert!(
plan.where_terms.is_empty(),
"precondition: no WHERE → empty where_terms"
);
reorder_predicates(&mut plan);
assert!(
plan.where_terms.is_empty(),
"where_terms must remain empty after reorder"
);
}
#[test]
fn pred_cost_orders_cheap_first() {
assert!(pred_cost_rank(PredCost::Type) < pred_cost_rank(PredCost::Scalar));
assert!(pred_cost_rank(PredCost::Scalar) < pred_cost_rank(PredCost::Str));
assert!(pred_cost_rank(PredCost::Str) < pred_cost_rank(PredCost::Ref));
}
#[test]
fn pred_cost_type_is_zero() {
assert_eq!(pred_cost_rank(PredCost::Type), 0);
}
#[test]
fn schema_stats_count_of_defaults_zero() {
let stats = SchemaStats::default();
assert_eq!(stats.count_of("java.lang.String"), 0);
assert_eq!(stats.count_of("com.example.Foo"), 0);
assert_eq!(stats.count_of(""), 0);
}
#[test]
fn schema_stats_count_of_returns_inserted() {
let mut stats = SchemaStats::default();
stats
.instance_counts
.insert("java.lang.String".to_string(), 42);
assert_eq!(stats.count_of("java.lang.String"), 42);
assert_eq!(stats.count_of("java.lang.Object"), 0);
}
#[test]
fn schema_stats_default_is_empty() {
assert!(SchemaStats::default().instance_counts.is_empty());
}
#[test]
fn limit_pushed_to_scan_when_safe() {
let mut plan = pq(&parse("SELECT @objectId FROM java.lang.String LIMIT 10").unwrap());
pushdown_limit(&mut plan);
assert_eq!(
plan.scan_limit,
Some(10),
"scan_limit must equal limit when pushdown is safe"
);
}
#[test]
fn limit_not_pushed_with_order_by() {
let mut plan = pq(&parse(
"SELECT @objectId FROM java.lang.String ORDER BY @retainedHeapSize LIMIT 10",
)
.unwrap());
pushdown_limit(&mut plan);
assert_eq!(
plan.scan_limit, None,
"ORDER BY @retainedHeapSize must block limit pushdown (order_sensitive + late_ops)"
);
}
#[test]
fn limit_not_pushed_with_scalar_order_by() {
let mut plan = pq(&parse(
"SELECT @objectId FROM java.lang.String ORDER BY @usedHeapSize LIMIT 10",
)
.unwrap());
assert!(
plan.late_ops.is_empty(),
"precondition: ORDER BY @usedHeapSize must not produce late ops, got {:?}",
plan.late_ops
);
assert!(
plan.order_sensitive,
"precondition: ORDER BY must set order_sensitive"
);
pushdown_limit(&mut plan);
assert_eq!(
plan.scan_limit, None,
"order_sensitive alone (no late ops) must block limit pushdown"
);
}
#[test]
fn no_limit_means_no_scan_limit() {
let mut plan = pq(&parse("SELECT @objectId FROM java.lang.String").unwrap());
pushdown_limit(&mut plan);
assert_eq!(
plan.scan_limit, None,
"absent LIMIT must leave scan_limit None"
);
}
#[test]
fn limit_not_pushed_with_late_ops() {
let mut plan =
pq(&parse("SELECT @retainedHeapSize FROM java.lang.String LIMIT 5").unwrap());
assert!(
!plan.late_ops.is_empty(),
"precondition: @retainedHeapSize in SELECT must produce late ops, got {:?}",
plan.late_ops
);
assert!(
!plan.order_sensitive,
"precondition: no ORDER BY → order_sensitive must be false"
);
pushdown_limit(&mut plan);
assert_eq!(
plan.scan_limit, None,
"non-empty late_ops must block limit pushdown"
);
}
#[test]
fn pushdown_is_idempotent() {
let mut plan = pq(&parse("SELECT @objectId FROM java.lang.String LIMIT 10").unwrap());
pushdown_limit(&mut plan);
assert_eq!(plan.scan_limit, Some(10), "first call must set scan_limit");
pushdown_limit(&mut plan);
assert_eq!(
plan.scan_limit,
Some(10),
"second call must not change scan_limit"
);
}
#[test]
fn projection_only_refpath_deferred() {
let query = parse("SELECT x.parent.name FROM Node x").unwrap();
let mut plan = pq(&query);
defer_projections(&mut plan, &query);
assert!(
!plan.deferred_projections.is_empty(),
"RefPath SELECT item must be marked deferrable; got {:?}",
plan.deferred_projections
);
assert_eq!(
plan.deferred_projections[0].select_index, 0,
"first (only) SELECT item is at index 0"
);
}
#[test]
fn retained_projection_deferred() {
let query = parse("SELECT @retainedHeapSize FROM java.lang.String").unwrap();
let mut plan = pq(&query);
defer_projections(&mut plan, &query);
assert!(
!plan.deferred_projections.is_empty(),
"@retainedHeapSize SELECT item must be marked deferrable; got {:?}",
plan.deferred_projections
);
assert_eq!(plan.deferred_projections[0].select_index, 0);
}
#[test]
fn plain_scalar_projection_not_deferred() {
let query = parse("SELECT @objectId FROM java.lang.String").unwrap();
let mut plan = pq(&query);
defer_projections(&mut plan, &query);
assert!(
plan.deferred_projections.is_empty(),
"@objectId is cheap — must NOT be marked deferrable; got {:?}",
plan.deferred_projections
);
}
#[test]
fn defer_is_idempotent() {
let query = parse("SELECT @retainedHeapSize FROM java.lang.String").unwrap();
let mut plan = pq(&query);
defer_projections(&mut plan, &query);
let first = plan.deferred_projections.clone();
defer_projections(&mut plan, &query);
let second = plan.deferred_projections.clone();
assert_eq!(
first, second,
"calling defer_projections twice must be idempotent"
);
}
#[test]
fn dead_retained_need_eliminated() {
let mut plan = pq(&parse("SELECT @usedHeapSize FROM java.lang.String").unwrap());
plan.late_ops.clear();
plan.needs.retained = true; eliminate_dead_needs(&mut plan);
assert!(
!plan.needs.retained,
"needs.retained must be cleared when no late op references it"
);
}
#[test]
fn live_retained_need_preserved() {
let mut plan = pq(&parse("SELECT @retainedHeapSize FROM java.lang.String").unwrap());
assert!(
plan.late_ops
.iter()
.any(|op| matches!(op, StageOp::JoinRetained)),
"precondition: @retainedHeapSize SELECT must produce JoinRetained, got {:?}",
plan.late_ops
);
assert!(
plan.needs.retained,
"precondition: needs.retained must be set by planner"
);
eliminate_dead_needs(&mut plan);
assert!(
plan.needs.retained,
"needs.retained must stay true when JoinRetained late op is present"
);
}
#[test]
fn eliminate_does_not_touch_scan_needs() {
let mut plan = pq(&parse(
"SELECT * FROM C s WHERE @displayName = \"foo\" \
AND count > 1 AND s INSTANCEOF java.lang.Object",
)
.unwrap());
assert!(plan.needs.instance_string, "precondition: instance_string");
assert!(plan.needs.instance_scalar, "precondition: instance_scalar");
assert!(plan.needs.runtime_type, "precondition: runtime_type");
let before_scalar = plan.needs.instance_scalar;
let before_string = plan.needs.instance_string;
let before_rt = plan.needs.runtime_type;
eliminate_dead_needs(&mut plan);
assert_eq!(
plan.needs.instance_scalar, before_scalar,
"instance_scalar must not change"
);
assert_eq!(
plan.needs.instance_string, before_string,
"instance_string must not change"
);
assert_eq!(
plan.needs.runtime_type, before_rt,
"runtime_type must not change"
);
}
#[test]
fn eliminate_is_idempotent() {
let mut plan = pq(&parse("SELECT @retainedHeapSize FROM java.lang.String").unwrap());
eliminate_dead_needs(&mut plan);
let needs_after_first = plan.needs.clone();
eliminate_dead_needs(&mut plan);
assert_eq!(
plan.needs, needs_after_first,
"eliminate_dead_needs must be idempotent"
);
}
#[test]
fn optimize_is_idempotent_and_composes() {
let src = "SELECT @objectId FROM java.lang.String LIMIT 3";
let q = parse(src).unwrap();
let plan = pq(&q);
let once = optimize(plan.clone(), &q, &SchemaStats::default());
let twice = optimize(once.clone(), &q, &SchemaStats::default());
assert_eq!(once, twice, "optimize must be idempotent");
assert_eq!(once.scan_limit, Some(3), "scan_limit must be pushed down");
}
#[test]
fn optimize_reorders_and_pushes_limit() {
let src = "SELECT @objectId FROM java.lang.String LIMIT 5";
let q = parse(src).unwrap();
let mut plan = pq(&q);
plan.where_terms = vec![
scalar_conjunct("d", PredCost::Ref),
scalar_conjunct("c", PredCost::Str),
scalar_conjunct("b", PredCost::Scalar),
scalar_conjunct("a", PredCost::Type),
];
let optimized = optimize(plan, &q, &SchemaStats::default());
let ranks: Vec<u8> = optimized
.where_terms
.iter()
.map(|c| pred_cost_rank(c.cost))
.collect();
assert!(
ranks.windows(2).all(|w| w[0] <= w[1]),
"where_terms must be sorted cheapest-first after optimize, got ranks: {:?}",
ranks
);
assert_eq!(
optimized.scan_limit,
Some(5),
"scan_limit must be pushed down by optimize"
);
}
#[test]
fn optimize_recurses_into_union_branches() {
let src =
"SELECT @objectId FROM java.lang.String UNION SELECT @objectId FROM java.lang.Object";
let q = parse(src).unwrap();
let plan = pq(&q);
assert_eq!(
plan.union_branches.len(),
1,
"precondition: one union branch"
);
let once = optimize(plan.clone(), &q, &SchemaStats::default());
let twice = optimize(once.clone(), &q, &SchemaStats::default());
assert_eq!(once, twice, "optimize must be idempotent for UNION queries");
assert_eq!(
once.union_branches.len(),
1,
"union branch must be preserved after optimize"
);
}
#[test]
fn optimize_default_queryplan_constructs() {
let d = QueryPlan::default();
assert_eq!(
d.kind,
StageKind::SingleScan,
"default kind must be SingleScan"
);
assert_eq!(
d.carry,
CarryLayout::IndexOnly,
"default carry must be IndexOnly"
);
assert_eq!(d.finalize_at, Phase::P1, "default finalize_at must be P1");
assert!(
d.where_terms.is_empty(),
"default where_terms must be empty"
);
assert!(d.late_ops.is_empty(), "default late_ops must be empty");
assert!(
d.union_branches.is_empty(),
"default union_branches must be empty"
);
assert!(
d.in_subplans.is_empty(),
"default in_subplans must be empty"
);
assert!(
d.deferred_projections.is_empty(),
"default deferred_projections must be empty"
);
assert!(
d.from_subplan.is_none(),
"default from_subplan must be None"
);
assert!(d.limit.is_none(), "default limit must be None");
assert!(d.scan_limit.is_none(), "default scan_limit must be None");
assert!(!d.order_sensitive, "default order_sensitive must be false");
assert_eq!(d.select_arity, 0, "default select_arity must be 0");
}
#[test]
fn order_by_selectivity_is_noop() {
let q = parse("SELECT @objectId FROM java.lang.String").unwrap();
let mut plan = pq(&q);
let snapshot = plan.clone();
order_by_selectivity(&mut plan, &SchemaStats::default());
assert_eq!(
plan, snapshot,
"order_by_selectivity must not change the plan"
);
}
#[test]
fn optimize_empty_where_and_no_limit() {
let src = "SELECT @objectId FROM java.lang.String";
let q = parse(src).unwrap();
let plan = pq(&q);
let once = optimize(plan.clone(), &q, &SchemaStats::default());
assert!(
once.where_terms.is_empty(),
"empty WHERE must remain empty after optimize"
);
assert_eq!(
once.scan_limit, None,
"absent LIMIT must leave scan_limit None after optimize"
);
let twice = optimize(once.clone(), &q, &SchemaStats::default());
assert_eq!(
once, twice,
"optimize must be idempotent on a simple no-WHERE no-LIMIT plan"
);
}
#[test]
fn narrow_carry_indexonly_is_noop() {
let mut plan = pq(&parse("SELECT @objectId FROM java.lang.String").unwrap());
assert!(
matches!(plan.carry, CarryLayout::IndexOnly),
"precondition: default carry is IndexOnly"
);
narrow_carry(&mut plan);
assert!(
matches!(plan.carry, CarryLayout::IndexOnly),
"narrow_carry must leave IndexOnly unchanged"
);
}
#[test]
fn narrow_carry_empty_scalars_downgrades() {
let mut plan = pq(&parse("SELECT @objectId FROM java.lang.String").unwrap());
plan.carry = CarryLayout::IndexPlusScalars { widths: vec![] };
narrow_carry(&mut plan);
assert!(
matches!(plan.carry, CarryLayout::IndexOnly),
"IndexPlusScalars with empty widths must downgrade to IndexOnly"
);
}
}