use arrow::datatypes::DataType;
use datafusion::datasource::memory::MemorySourceConfig;
use datafusion::datasource::source::DataSourceExec;
use datafusion::physical_plan::ExecutionPlan;
use crate::agg_table as table;
use crate::exec::{AggKind, MetalOp};
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct Shape {
pub families: Vec<&'static str>,
pub keys: &'static str,
pub key_class: &'static str,
pub source: &'static str,
}
impl std::fmt::Display for Shape {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"{} over {} {} key{} ({} input)",
self.families.join(" + "),
self.keys,
self.key_class,
if self.keys == "1" { "" } else { "s" },
self.source
)
}
}
fn key_class(t: &DataType) -> &'static str {
match t {
DataType::Int8 | DataType::Int16 | DataType::Int32 | DataType::UInt8 | DataType::UInt16 => "i32",
DataType::Int64 | DataType::UInt32 | DataType::UInt64 => "i64",
DataType::Float32 | DataType::Float64 => "float",
_ => "string",
}
}
fn is_int(t: &DataType) -> bool {
t.is_integer()
}
fn source_class(input: &dyn ExecutionPlan) -> &'static str {
match input.downcast_ref::<DataSourceExec>() {
Some(d) if d.data_source().downcast_ref::<MemorySourceConfig>().is_some() => "memory",
_ => "other",
}
}
pub(crate) fn shape(op: &MetalOp, input: &dyn ExecutionPlan) -> Option<Shape> {
let MetalOp::Aggregate { keys, aggs } = op else { return None };
let schema = input.schema();
let rank = |c: &str| ["i32", "i64", "float", "string"].iter().position(|x| *x == c).unwrap_or(3);
let mut kc = "i32";
for &k in keys {
let c = key_class(schema.field(k).data_type());
if rank(c) > rank(kc) {
kc = c;
}
}
let mut families: Vec<&'static str> = aggs
.iter()
.map(|a| {
let t = a.column.map(|c| schema.field(c).data_type().clone());
match (a.kind, t) {
(AggKind::CountAll | AggKind::Count, _) => "count",
(AggKind::Sum | AggKind::Mean, Some(t)) if is_int(&t) => "sum_avg_int",
(AggKind::Sum | AggKind::Mean, Some(DataType::Float64)) => "sum_avg_f64",
(AggKind::Sum | AggKind::Mean, _) => "sum_avg_f32",
(AggKind::Min | AggKind::Max, Some(t)) if is_int(&t) => "minmax_int",
(AggKind::Min | AggKind::Max, _) => "minmax_f64",
}
})
.collect();
if families.is_empty() {
families.push("distinct");
}
families.sort_unstable();
families.dedup();
Some(Shape { families, keys: if keys.len() == 1 { "1" } else { "2+" }, key_class: kc, source: source_class(input) })
}
pub(crate) fn bucket(groups: u64, rows: u64) -> Option<&'static str> {
if groups == 0 {
return None;
}
if groups.saturating_mul(table::NEAR_ROWS) >= rows {
return Some(table::ROWS_BUCKET);
}
table::BUCKETS.iter().find(|(_, lo, hi)| *lo <= groups && groups <= *hi).map(|b| b.0)
}
fn row(family: &str, s: &Shape, bucket: &str) -> Option<&'static table::Row> {
table::TABLE
.iter()
.find(|r| {
r.family == family && r.keys == s.keys && r.key_class == s.key_class && r.source == s.source && r.bucket == bucket
})
}
fn family_takes(family: &str, s: &Shape, bucket: &str, rows: u64) -> Result<u64, String> {
let Some(r) = row(family, s, bucket) else {
return Err(format!(
"{family} over {} {} key(s), {} input, at the {bucket} bucket was not measured",
s.keys, s.key_class, s.source
));
};
let ratios = r.ratios.iter().map(|(n, x)| format!("{}: {x:.2}x", rows_text(*n))).collect::<Vec<_>>().join(", ");
match r.min_rows {
None => Err(format!(
"{family} at the {bucket} bucket is not taken at any measured size (needs {:.2}x at two or more sizes; {ratios})",
table::MIN_RATIO * table::HEADROOM
)),
Some(m) if rows < m => Err(format!("{family} at the {bucket} bucket is taken from {} rows ({ratios})", rows_text(m))),
Some(_) if r.max_rows.is_some_and(|x| rows > x) => Err(format!(
"{family} at the {bucket} bucket is taken up to {} rows only, the largest size measured, its ratio falling there ({ratios})",
rows_text(r.max_rows.unwrap_or_default())
)),
Some(m) => Ok(m),
}
}
pub(crate) fn rows_text(n: u64) -> String {
if n >= 1_000_000 && n.is_multiple_of(1_000_000) {
format!("{}M", n / 1_000_000)
} else if n >= 1_000 && n.is_multiple_of(1_000) {
format!("{}k", n / 1_000)
} else {
n.to_string()
}
}
pub(crate) fn takes(s: &Shape, bucket: Option<&str>, rows: u64) -> Result<String, String> {
let Some(b) = bucket else {
return Err("no bucket holds that group count at this row count".into());
};
let mut from = 0;
for f in &s.families {
from = from.max(family_takes(f, s, b, rows)?);
}
Ok(format!("{} at the {b} bucket is taken from {} rows", s.families.join(" + "), rows_text(from)))
}
fn regions(rows: u64) -> Vec<(u64, u64, Option<&'static str>)> {
let edge = rows.div_ceil(table::NEAR_ROWS).max(1);
let mut out = Vec::new();
let mut top = 0;
for &(name, lo, hi) in table::BUCKETS {
if lo >= edge {
break;
}
top = hi.min(edge - 1);
out.push((lo, top, Some(name)));
}
if top + 1 < edge {
out.push((top + 1, edge - 1, None));
}
out.push((edge, u64::MAX, Some(table::ROWS_BUCKET)));
out
}
pub(crate) fn settled_for(s: &Shape, rows: u64) -> impl Fn(u64, u64) -> bool {
let regions: Vec<(u64, u64, bool)> =
regions(rows).into_iter().map(|(lo, hi, b)| (lo, hi, takes(s, b, rows).is_ok())).collect();
move |lo, hi| {
let mut seen: Option<bool> = None;
for &(rlo, rhi, t) in ®ions {
if rlo <= hi && rhi >= lo {
if seen.is_some_and(|x| x != t) {
return false;
}
seen = Some(t);
}
}
true
}
}
pub(crate) fn any_bucket(s: &Shape, rows: u64) -> Result<String, String> {
let taken: Vec<&str> = regions(rows).into_iter().filter_map(|(_, _, b)| b).filter(|b| takes(s, Some(b), rows).is_ok()).collect();
if taken.is_empty() {
let smallest = smallest_rows(s);
return Err(match smallest {
Some(m) => format!(
"{s}: the measured table takes no group count at {} rows (the first take is at {} rows; {})",
rows_text(rows),
rows_text(m),
table::SOURCE
),
None => format!("{s}: the measured table takes it at no group count and size ({})", table::SOURCE),
});
}
Ok(format!("{s}: the measured table takes the {} bucket(s) at {} rows; MetalExec decides from a group-count estimate", taken.join(", "), rows_text(rows)))
}
fn join_key_class(left: &dyn ExecutionPlan, keys: &[usize]) -> &'static str {
let schema = left.schema();
let rank = |c: &str| ["i32", "i64", "string"].iter().position(|x| *x == c).unwrap_or(2);
let mut kc = "i32";
for &k in keys {
let c = match schema.field(k).data_type() {
DataType::Int32 => "i32",
DataType::Int64 => "i64",
_ => "string",
};
if rank(c) > rank(kc) {
kc = c;
}
}
kc
}
pub(crate) fn join_takes(
op: &MetalOp,
left: &dyn ExecutionPlan,
right: &dyn ExecutionPlan,
build: u64,
probe: u64,
) -> Result<String, String> {
use crate::join_table as jt;
let MetalOp::Join { how, left_keys, .. } = op else {
return Err("not a join".into());
};
let how = match how {
crate::exec::JoinHow::Inner => "inner",
crate::exec::JoinHow::Left => "left",
crate::exec::JoinHow::Right => "right",
};
let key_class = join_key_class(left, left_keys);
let source = if source_class(left) == "memory" && source_class(right) == "memory" { "memory" } else { "other" };
let shape = format!("{how} join on {key_class} keys ({source} inputs)");
let Some(&(bucket, _, _)) = jt::BUILD_BUCKETS.iter().find(|(_, lo, hi)| *lo <= build && build <= *hi) else {
return Err(format!("{shape}: no build-side bucket of the measured join table holds {} rows ({})", rows_text(build), jt::SOURCE));
};
let Some(r) = jt::TABLE
.iter()
.find(|r| r.how == how && r.key_class == key_class && r.source == source && r.build == bucket)
else {
return Err(format!("{shape}, {bucket}-row build bucket: not measured ({})", jt::SOURCE));
};
match r.min_probe_rows {
Some(m) if probe >= m => Ok(format!(
"{shape}, {bucket}-row build bucket: the measured join table takes it from {} probe rows ({})",
rows_text(m),
r.reason
)),
Some(m) => Err(format!(
"{shape}, {bucket}-row build bucket: the measured join table takes it from {} probe rows ({})",
rows_text(m),
r.reason
)),
None => Err(format!("{shape}, {bucket}-row build bucket: not taken ({}; {})", r.reason, jt::SOURCE)),
}
}
fn smallest_rows(s: &Shape) -> Option<u64> {
let mut best: Option<u64> = None;
let buckets = table::BUCKETS.iter().map(|b| b.0).chain(std::iter::once(table::ROWS_BUCKET));
for b in buckets {
let mut need = Some(0u64);
for f in &s.families {
need = match (need, row(f, s, b).and_then(|r| r.min_rows)) {
(Some(a), Some(m)) => Some(a.max(m)),
_ => None,
};
}
if let Some(n) = need {
best = Some(best.map_or(n, |x: u64| x.min(n)));
}
}
best
}