use std::cmp::Ordering::Equal;
use reifydb_core::{
error::diagnostic::query,
sort::{
SortDirection,
SortDirection::{Asc, Desc},
SortKey,
},
value::column::{buffer::ColumnBuffer, columns::Columns, headers::ColumnHeaders},
};
use reifydb_extension::transform::{Transform, context::TransformContext};
use reifydb_transaction::transaction::Transaction;
use reifydb_value::{error, error::Error, reifydb_assertions};
use tracing::instrument;
use crate::{
Result,
vm::volcano::query::{QueryContext, QueryNode, charge_query_memory},
};
pub(crate) struct SortNode {
input: Box<dyn QueryNode>,
by: Vec<SortKey>,
initialized: Option<()>,
}
impl SortNode {
pub(crate) fn new(input: Box<dyn QueryNode>, by: Vec<SortKey>) -> Self {
Self {
input,
by,
initialized: None,
}
}
#[instrument(level = "trace", skip_all, name = "volcano::sort::collect")]
fn collect_input<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &mut QueryContext) -> Result<Option<Columns>> {
let mut columns_opt: Option<Columns> = None;
let mut charged = 0usize;
while let Some(columns) = self.input.next(rx, ctx)? {
if let Some(existing_columns) = &mut columns_opt {
existing_columns.system.extend(&columns.system)?;
for (i, col) in columns.columns.iter().enumerate() {
existing_columns[i].extend(col.clone())?;
}
} else {
columns_opt = Some(columns);
}
if let Some(acc) = &columns_opt {
charge_query_memory(&ctx.memory, &mut charged, acc)?;
}
}
Ok(columns_opt)
}
}
impl QueryNode for SortNode {
#[instrument(level = "trace", skip_all, name = "volcano::sort::initialize")]
fn initialize<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &QueryContext) -> Result<()> {
self.input.initialize(rx, ctx)?;
self.initialized = Some(());
Ok(())
}
#[instrument(level = "trace", skip_all, name = "volcano::sort::next")]
fn next<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &mut QueryContext) -> Result<Option<Columns>> {
reifydb_assertions! {
assert!(self.initialized.is_some(), "SortNode::next() called before initialize()");
}
let columns_opt = self.collect_input(rx, ctx)?;
let columns = match columns_opt {
Some(f) => f,
None => return Ok(None),
};
let transform_ctx = TransformContext {
routines: &ctx.services.routines,
runtime_context: &ctx.services.runtime_context,
params: &ctx.params,
};
Ok(Some(self.apply(&transform_ctx, columns)?))
}
fn headers(&self) -> Option<ColumnHeaders> {
self.input.headers()
}
}
impl Transform for SortNode {
fn apply(&self, _ctx: &TransformContext, mut columns: Columns) -> Result<Columns> {
let key_refs =
self.by.iter()
.map(|key| {
let name = key.column.fragment();
if let Some(data) = columns.system_column(name) {
return Ok::<_, Error>((data, key.direction.clone()));
}
let col = columns
.iter()
.find(|c| c.name() == name)
.ok_or_else(|| error!(query::column_not_found(key.column.clone())))?;
Ok((col.data().clone(), key.direction.clone()))
})
.collect::<Result<Vec<_>>>()?;
let indices = Self::rank_rows(&key_refs, columns.row_count());
Self::permute(&mut columns, &indices);
Ok(columns)
}
}
impl SortNode {
#[instrument(level = "trace", skip_all, name = "volcano::sort::rank")]
fn rank_rows(key_refs: &[(ColumnBuffer, SortDirection)], row_count: usize) -> Vec<usize> {
let mut indices: Vec<usize> = (0..row_count).collect();
indices.sort_unstable_by(|&l, &r| {
for (col, dir) in key_refs {
let vl = col.get_value(l);
let vr = col.get_value(r);
let ord = vl.partial_cmp(&vr).unwrap_or(Equal);
let ord = match dir {
Asc => ord,
Desc => ord.reverse(),
};
if ord != Equal {
return ord;
}
}
Equal
});
indices
}
#[instrument(level = "trace", skip_all, name = "volcano::sort::permute")]
fn permute(columns: &mut Columns, indices: &[usize]) {
columns.system.permute_in_place(indices);
let cols = &mut columns.columns;
for col in cols.iter_mut() {
col.reorder(indices);
}
}
}