use crate::{Partitioning, PhysicalPlan};
use super::{AqeRule, RuntimeStats, StreamingAqeGuard};
const DEFAULT_TARGET_PARTITION_BYTES: u64 = krishiv_common::partition::TARGET_BYTES_PER_PARTITION;
pub struct AutoPartitionRule {
target_partition_bytes: u64,
max_buckets: u32,
}
impl AutoPartitionRule {
pub fn new(max_buckets: u32) -> Self {
Self {
target_partition_bytes: DEFAULT_TARGET_PARTITION_BYTES,
max_buckets,
}
}
#[must_use]
pub fn with_target_partition_bytes(mut self, bytes: u64) -> Self {
self.target_partition_bytes = bytes;
self
}
}
impl AqeRule for AutoPartitionRule {
fn name(&self) -> &str {
"auto-partition"
}
fn apply(&self, plan: &PhysicalPlan, stats: &[RuntimeStats]) -> Option<PhysicalPlan> {
if let Some(override_buckets) = plan.shuffle_partitions() {
return self.apply_override(plan, override_buckets);
}
if stats.is_empty() || StreamingAqeGuard::plan_is_streaming(plan) {
return None;
}
let total_bytes: u64 = stats
.iter()
.map(|s| {
if s.serialized_bytes > 0 {
s.serialized_bytes
} else {
s.memory_bytes
}
})
.sum();
if total_bytes == 0 {
return None;
}
let target = krishiv_common::partition::recommend_buckets(
total_bytes,
1,
self.max_buckets,
self.target_partition_bytes,
);
self.stamp_target(plan, target)
}
}
impl AutoPartitionRule {
fn apply_override(&self, plan: &PhysicalPlan, target: u32) -> Option<PhysicalPlan> {
if StreamingAqeGuard::plan_is_streaming(plan) {
return None;
}
let target = target.max(1);
self.stamp_target(plan, target)
}
fn stamp_target(&self, plan: &PhysicalPlan, target: u32) -> Option<PhysicalPlan> {
let mut changed = false;
for node in plan.nodes() {
match node.partitioning() {
Partitioning::Hash { buckets, .. } | Partitioning::RoundRobin { buckets, .. }
if *buckets != target =>
{
changed = true;
}
_ => {}
}
}
if !changed {
return None;
}
let mut plan = plan.clone();
for node in plan.nodes_mut() {
let old = node.partitioning().clone();
match old {
Partitioning::Hash { ref keys, buckets } if buckets != target => {
node.set_partitioning(Partitioning::Hash {
keys: keys.clone(),
buckets: target,
});
}
Partitioning::RoundRobin { buckets } if buckets != target => {
node.set_partitioning(Partitioning::RoundRobin { buckets: target });
}
_ => {}
}
}
tracing::debug!(rule = "auto-partition", target, "AutoPartitionRule applied");
Some(plan)
}
}