use std::{mem, sync::Arc};
use reifydb_core::value::column::{ColumnWithName, buffer::ColumnBuffer, columns::Columns, headers::ColumnHeaders};
use reifydb_evaluate::expression::{
compile::{CompiledExpr, compile_expression},
context::{CompileContext, EvalContext},
};
use reifydb_extension::transform::{Transform, context::TransformContext};
use reifydb_rql::expression::Expression;
use reifydb_transaction::transaction::Transaction;
use reifydb_value::{reifydb_assertions, util::bitvec::BitVec};
use tracing::instrument;
use super::NoopNode;
use crate::{
Result,
vm::volcano::{
query::{QueryContext, QueryNode, eval_context_from_transform},
udf::{UdfEvalNode, strip_udf_columns},
},
};
pub(crate) struct FilterNode {
input: Box<dyn QueryNode>,
expressions: Vec<Expression>,
udf_names: Vec<String>,
context: Option<(Arc<QueryContext>, Vec<CompiledExpr>)>,
}
impl FilterNode {
pub fn new(input: Box<dyn QueryNode>, expressions: Vec<Expression>) -> Self {
Self {
input,
expressions,
udf_names: Vec::new(),
context: None,
}
}
#[instrument(level = "trace", skip_all, name = "volcano::filter::eval")]
fn eval_predicate(
session: &EvalContext,
compiled: &CompiledExpr,
columns: &Columns,
row_count: usize,
) -> Result<ColumnWithName> {
let exec_ctx = session.with_eval(columns.clone(), row_count);
compiled.execute(&exec_ctx)
}
#[instrument(level = "trace", skip_all, name = "volcano::filter::mask")]
fn build_mask(result: &ColumnBuffer, row_count: usize) -> BitVec {
match result {
ColumnBuffer::Bool(container) => {
let mut mask = BitVec::repeat(row_count, false);
for i in 0..row_count {
if i < container.len() {
let valid = container.is_defined(i);
let filter_result = container.data().get(i);
mask.set(i, valid & filter_result);
}
}
mask
}
ColumnBuffer::Option {
inner,
bitvec,
} => match inner.as_ref() {
ColumnBuffer::Bool(container) => {
let mut mask = BitVec::repeat(row_count, false);
for i in 0..row_count {
let defined = i < bitvec.len() && bitvec.get(i);
let valid = defined && container.is_defined(i);
let value = valid && container.data().get(i);
mask.set(i, value);
}
mask
}
_ => panic!("filter expression must evaluate to a boolean column"),
},
_ => panic!("filter expression must evaluate to a boolean column"),
}
}
#[instrument(level = "trace", skip_all, name = "volcano::filter::compact")]
fn compact(columns: &mut Columns, mask: &BitVec) -> Result<()> {
columns.filter(mask)
}
}
impl QueryNode for FilterNode {
#[instrument(level = "trace", skip_all, name = "volcano::filter::initialize")]
fn initialize<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &QueryContext) -> Result<()> {
let (input, expressions, udf_names) = UdfEvalNode::wrap_if_needed(
mem::replace(&mut self.input, Box::new(NoopNode)),
&self.expressions,
&ctx.symbols,
);
self.input = input;
self.expressions = expressions;
self.udf_names = udf_names;
let compile_ctx = CompileContext {
symbols: &ctx.symbols,
};
let compiled = self
.expressions
.iter()
.map(|e| compile_expression(&compile_ctx, e).expect("compile"))
.collect();
self.context = Some((Arc::new(ctx.clone()), compiled));
self.input.initialize(rx, ctx)?;
Ok(())
}
#[instrument(level = "trace", skip_all, name = "volcano::filter::next")]
fn next<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &mut QueryContext) -> Result<Option<Columns>> {
reifydb_assertions! {
assert!(self.context.is_some(), "FilterNode::next() called before initialize()");
}
let (stored_ctx, _) = self.context.as_ref().unwrap();
let stored_ctx = stored_ctx.clone();
loop {
match self.input.next(rx, ctx)? {
Some(columns) => {
let transform_ctx = TransformContext {
routines: &ctx.services.routines,
runtime_context: &stored_ctx.services.runtime_context,
params: &stored_ctx.params,
};
let mut columns = self.apply(&transform_ctx, columns)?;
if columns.row_count() > 0 {
strip_udf_columns(&mut columns, &self.udf_names);
return Ok(Some(columns));
}
}
None => return Ok(None),
}
}
}
fn headers(&self) -> Option<ColumnHeaders> {
self.input.headers()
}
}
impl Transform for FilterNode {
fn apply(&self, ctx: &TransformContext, input: Columns) -> Result<Columns> {
let (stored_ctx, compiled) =
self.context.as_ref().expect("FilterNode::apply() called before initialize()");
let session = eval_context_from_transform(ctx, stored_ctx);
let mut columns = input;
let mut row_count = columns.row_count();
for compiled_expr in compiled {
if row_count == 0 {
break;
}
let result = Self::eval_predicate(&session, compiled_expr, &columns, row_count)?;
let filter_mask = Self::build_mask(result.data(), row_count);
Self::compact(&mut columns, &filter_mask)?;
row_count = columns.row_count();
}
Ok(columns)
}
}