use super::frame::{BoundKind, CurrentRow, FrameCursor, FrameSpec, Membership};
use super::partition::PartitionRows;
use crate::aggregation::{
aggregate_value_with_args, instantiate_aggregate_accumulators, observe_aggregate,
AggregateAccumulator, AggregateAccumulatorTemplate,
};
use uqa_core::Value;
use uqa_sql::ast::FrameExclusion;
use uqa_sql::{SQLError, ScalarExpr};
pub(super) struct WindowAggregate {
name: String,
args: Vec<ScalarExpr>,
filter: Option<ScalarExpr>,
template: AggregateAccumulatorTemplate,
budget_bytes: usize,
accumulator: AggregateAccumulator,
aggregated_base: i64,
aggregated_upto: i64,
result: Value,
}
impl WindowAggregate {
pub(super) fn new(
(name, args, filter): (&str, &[ScalarExpr], Option<&ScalarExpr>),
template: AggregateAccumulatorTemplate,
budget_bytes: usize,
) -> Self {
let accumulator = instantiate(&template, budget_bytes);
Self {
name: name.to_string(),
args: args.to_vec(),
filter: filter.cloned(),
template,
budget_bytes,
accumulator,
aggregated_base: 0,
aggregated_upto: 0,
result: Value::Null,
}
}
pub(super) fn begin_partition(&mut self) {
self.aggregated_base = 0;
self.aggregated_upto = 0;
self.result = Value::Null;
}
pub(super) fn value(
&mut self,
frame: &FrameSpec,
cursor: &mut FrameCursor,
current: &mut CurrentRow,
rows: &mut PartitionRows<'_>,
) -> Result<Value, SQLError> {
let head = cursor.head(frame, current, rows)?;
if head < self.aggregated_base {
return Err(SQLError::Internal(
"window frame head moved backward".into(),
));
}
if self.aggregated_base == head
&& matches!(
frame.end,
BoundKind::UnboundedFollowing | BoundKind::CurrentRow
)
&& frame.exclusion == FrameExclusion::NoOthers
&& self.aggregated_base <= current.position
&& self.aggregated_upto > current.position
{
return Ok(self.result.clone());
}
if current.position == 0
|| self.aggregated_base != head
|| frame.exclusion != FrameExclusion::NoOthers
|| self.aggregated_upto <= head
{
self.accumulator = instantiate(&self.template, self.budget_bytes);
self.aggregated_upto = head;
}
self.aggregated_base = head;
while self.aggregated_upto < rows.len() {
match cursor.membership(frame, current, rows, self.aggregated_upto)? {
Membership::After => break,
Membership::Outside => {}
Membership::Inside => {
let (name, args, filter, accumulator) =
(&self.name, &self.args, &self.filter, &mut self.accumulator);
rows.with_context(self.aggregated_upto, |context| {
if let Some(filter) = filter {
if !uqa_sql::expr::truthy(&crate::eval_scalar(filter, context)?) {
return Ok(());
}
}
observe_aggregate(accumulator, name, args, false, &[], context)
})?;
}
}
self.aggregated_upto += 1;
}
self.result = aggregate_value_with_args(
&self.name,
&self.accumulator,
&self.args,
rows.enum_labels(),
)?;
Ok(self.result.clone())
}
}
fn instantiate(
template: &AggregateAccumulatorTemplate,
budget_bytes: usize,
) -> AggregateAccumulator {
instantiate_aggregate_accumulators(std::slice::from_ref(template), budget_bytes)
.pop()
.expect("one template instantiates one accumulator")
}