use super::aggregates::WindowAggregate;
use super::frame::{CurrentRow, FrameCursor, FrameSpec};
use super::partition::PartitionRows;
use uqa_core::Value;
use uqa_sql::{SQLError, ScalarExpr};
pub(super) enum WindowFunction {
RowNumber,
Rank,
DenseRank,
PercentRank,
CumeDist,
Ntile(ScalarExpr),
Shift(Box<Shift>),
FirstValue(ScalarExpr),
LastValue(ScalarExpr),
NthValue {
target: ScalarExpr,
position: ScalarExpr,
},
Aggregate(Box<WindowAggregate>),
}
pub(super) struct Shift {
pub(super) forward: bool,
pub(super) target: ScalarExpr,
pub(super) offset: Option<ScalarExpr>,
pub(super) default: Option<ScalarExpr>,
}
#[derive(Default)]
struct NtileState {
bucket: i64,
rows_in_bucket: i64,
boundary: i64,
remainder: i64,
}
pub(super) struct WindowFunctionState {
function: WindowFunction,
frame: FrameSpec,
cursor: FrameCursor,
rank: i64,
ntile: NtileState,
}
impl WindowFunctionState {
pub(super) fn new(function: WindowFunction, frame: FrameSpec) -> Self {
Self {
function,
frame,
cursor: FrameCursor::new(),
rank: 0,
ntile: NtileState::default(),
}
}
pub(super) fn frame_mut(&mut self) -> &mut FrameSpec {
&mut self.frame
}
pub(super) fn begin_partition(&mut self) {
self.cursor = FrameCursor::new();
self.rank = 0;
self.ntile = NtileState::default();
if let WindowFunction::Aggregate(aggregate) = &mut self.function {
aggregate.begin_partition();
}
}
pub(super) fn advance(&mut self) {
self.cursor.invalidate();
}
pub(super) fn value(
&mut self,
current: &mut CurrentRow,
rows: &mut PartitionRows<'_>,
) -> Result<Value, SQLError> {
let position = current.position;
match &mut self.function {
WindowFunction::RowNumber => Ok(Value::Int(position + 1)),
WindowFunction::Rank => {
if rank_up(&mut self.rank, current, rows)? {
self.rank = position + 1;
}
Ok(Value::Int(self.rank))
}
WindowFunction::DenseRank => {
if rank_up(&mut self.rank, current, rows)? {
self.rank += 1;
}
Ok(Value::Int(self.rank))
}
WindowFunction::PercentRank => {
if rank_up(&mut self.rank, current, rows)? {
self.rank = position + 1;
}
let total = rows.len();
Ok(Value::Float(if total <= 1 {
0.0
} else {
(self.rank - 1) as f64 / (total - 1) as f64
}))
}
WindowFunction::CumeDist => {
let up = rank_up(&mut self.rank, current, rows)?;
if up || self.rank == 1 {
self.rank = position + 1;
let mut row = self.rank;
while row < rows.len() && rows.are_peers(row - 1, row)? {
self.rank += 1;
row += 1;
}
}
Ok(Value::Float(self.rank as f64 / rows.len() as f64))
}
WindowFunction::Ntile(argument) => ntile(&mut self.ntile, argument, current, rows),
WindowFunction::Shift(shift) => shift_value(shift, current, rows),
WindowFunction::FirstValue(target) => frame_value(
&self.frame,
&mut self.cursor,
target,
(true, 0),
current,
rows,
),
WindowFunction::LastValue(target) => frame_value(
&self.frame,
&mut self.cursor,
target,
(false, 0),
current,
rows,
),
WindowFunction::NthValue {
target,
position: nth,
} => {
let nth = match rows.evaluate(nth, position)? {
Value::Null => return Ok(Value::Null),
Value::Int(nth) => nth,
other => {
return Err(SQLError::Internal(format!(
"nth_value position {other:?} is not an integer"
)))
}
};
if nth <= 0 {
return Err(SQLError::Routine {
sqlstate: "22016".into(),
message: "argument of nth_value must be greater than zero".into(),
});
}
frame_value(
&self.frame,
&mut self.cursor,
target,
(true, nth - 1),
current,
rows,
)
}
WindowFunction::Aggregate(aggregate) => {
aggregate.value(&self.frame, &mut self.cursor, current, rows)
}
}
}
}
fn shift_value(
shift: &Shift,
current: &CurrentRow,
rows: &mut PartitionRows<'_>,
) -> Result<Value, SQLError> {
let position = current.position;
let offset = match &shift.offset {
Some(offset) => match rows.evaluate(offset, position)? {
Value::Null => return Ok(Value::Null),
Value::Int(offset) => offset,
other => {
return Err(SQLError::Internal(format!(
"window offset {other:?} is not an integer"
)))
}
},
None => 1,
};
let target_position = if shift.forward {
position.checked_add(offset)
} else {
position.checked_sub(offset)
};
match target_position.filter(|target| (0..rows.len()).contains(target)) {
Some(target_position) => rows.evaluate(&shift.target, target_position),
None => shift
.default
.as_ref()
.map_or(Ok(Value::Null), |default| rows.evaluate(default, position)),
}
}
fn frame_value(
frame: &FrameSpec,
cursor: &mut FrameCursor,
target: &ScalarExpr,
(from_head, offset): (bool, i64),
current: &mut CurrentRow,
rows: &mut PartitionRows<'_>,
) -> Result<Value, SQLError> {
cursor
.seek(frame, current, rows, from_head, offset)?
.map_or(Ok(Value::Null), |found| rows.evaluate(target, found))
}
fn rank_up(
rank: &mut i64,
current: &CurrentRow,
rows: &mut PartitionRows<'_>,
) -> Result<bool, SQLError> {
if *rank == 0 {
*rank = 1;
return Ok(false);
}
Ok(!rows.are_peers(current.position - 1, current.position)?)
}
fn ntile(
state: &mut NtileState,
argument: &ScalarExpr,
current: &CurrentRow,
rows: &mut PartitionRows<'_>,
) -> Result<Value, SQLError> {
if state.bucket == 0 {
let buckets = match rows.evaluate(argument, current.position)? {
Value::Null => return Ok(Value::Null),
Value::Int(buckets) => buckets,
other => {
return Err(SQLError::Internal(format!(
"ntile bucket count {other:?} is not an integer"
)))
}
};
if buckets <= 0 {
return Err(SQLError::Routine {
sqlstate: "22014".into(),
message: "argument of ntile must be greater than zero".into(),
});
}
let total = rows.len();
state.bucket = 1;
state.rows_in_bucket = 0;
state.boundary = total / buckets;
if state.boundary <= 0 {
state.boundary = 1;
} else {
state.remainder = total % buckets;
if state.remainder != 0 {
state.boundary += 1;
}
}
}
state.rows_in_bucket += 1;
if state.boundary < state.rows_in_bucket {
if state.remainder != 0 && state.bucket == state.remainder {
state.remainder = 0;
state.boundary -= 1;
}
state.bucket += 1;
state.rows_in_bucket = 1;
}
Ok(Value::Int(state.bucket))
}