use std::collections::HashMap;
use std::rc::Rc;
use std::sync::{Arc, Mutex};
use std::time::Instant;
use super::ast::{BinaryOp, Expr, LiteralValue, UnaryOp};
use super::ops;
use super::vectorized::{load_column_values, ValueVec, VectorizedEvaluator};
use crate::core::error::{Error, Result};
use crate::dataframe::base::DataFrame;
use crate::lock_safe;
#[derive(Debug, Clone, Default)]
pub struct JitQueryStats {
pub compilations: u64,
pub jit_executions: u64,
pub native_executions: u64,
pub compilation_time_ns: u64,
pub jit_execution_time_ns: u64,
pub native_execution_time_ns: u64,
pub vectorized_executions: u64,
pub vectorized_execution_time_ns: u64,
pub row_by_row_executions: u64,
pub row_by_row_execution_time_ns: u64,
}
impl JitQueryStats {
pub fn record_compilation(&mut self, duration_ns: u64) {
self.compilations += 1;
self.compilation_time_ns += duration_ns;
}
pub fn record_jit_execution(&mut self, duration_ns: u64) {
self.jit_executions += 1;
self.jit_execution_time_ns += duration_ns;
}
pub fn record_native_execution(&mut self, duration_ns: u64) {
self.native_executions += 1;
self.native_execution_time_ns += duration_ns;
}
pub fn record_vectorized_execution(&mut self, duration_ns: u64) {
self.vectorized_executions += 1;
self.vectorized_execution_time_ns += duration_ns;
self.record_native_execution(duration_ns);
}
pub fn record_row_by_row_execution(&mut self, duration_ns: u64) {
self.row_by_row_executions += 1;
self.row_by_row_execution_time_ns += duration_ns;
self.record_native_execution(duration_ns);
}
pub fn average_compilation_time_ns(&self) -> f64 {
if self.compilations > 0 {
self.compilation_time_ns as f64 / self.compilations as f64
} else {
0.0
}
}
pub fn vectorized_speedup_ratio(&self) -> f64 {
if self.vectorized_executions > 0 && self.row_by_row_executions > 0 {
let avg_row =
self.row_by_row_execution_time_ns as f64 / self.row_by_row_executions as f64;
let avg_vec =
self.vectorized_execution_time_ns as f64 / self.vectorized_executions as f64;
if avg_vec > 0.0 {
return avg_row / avg_vec;
}
}
1.0
}
pub fn jit_speedup_ratio(&self) -> f64 {
if self.jit_executions > 0 && self.native_executions > 0 {
let avg_native = self.native_execution_time_ns as f64 / self.native_executions as f64;
let avg_jit = self.jit_execution_time_ns as f64 / self.jit_executions as f64;
if avg_jit > 0.0 {
return avg_native / avg_jit;
}
}
1.0
}
}
#[derive(Clone)]
struct CompiledExpression {
#[allow(dead_code)] signature: String,
prepared: Option<Expr>,
execution_count: u64,
#[allow(dead_code)] last_execution: std::time::SystemTime,
}
pub struct QueryContext {
pub variables: HashMap<String, LiteralValue>,
pub functions: HashMap<String, Box<dyn Fn(&[f64]) -> f64 + Send + Sync>>,
compiled_expressions: Arc<Mutex<HashMap<String, CompiledExpression>>>,
jit_stats: Arc<Mutex<JitQueryStats>>,
jit_threshold: u64,
jit_enabled: bool,
}
impl std::fmt::Debug for QueryContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("QueryContext")
.field("variables", &self.variables)
.field("functions", &format!("{} functions", self.functions.len()))
.finish()
}
}
impl Default for QueryContext {
fn default() -> Self {
let mut context = Self {
variables: HashMap::new(),
functions: HashMap::new(),
compiled_expressions: Arc::new(Mutex::new(HashMap::new())),
jit_stats: Arc::new(Mutex::new(JitQueryStats::default())),
jit_threshold: 5, jit_enabled: true,
};
context.add_builtin_functions();
context
}
}
impl QueryContext {
pub fn new() -> Self {
Self::default()
}
pub fn with_jit_settings(jit_enabled: bool, jit_threshold: u64) -> Self {
let mut context = Self::default();
context.jit_enabled = jit_enabled;
context.jit_threshold = jit_threshold;
context
}
pub fn set_variable(&mut self, name: String, value: LiteralValue) {
self.variables.insert(name, value);
}
pub fn add_function<F>(&mut self, name: String, func: F)
where
F: Fn(&[f64]) -> f64 + Send + Sync + 'static,
{
self.functions.insert(name, Box::new(func));
}
pub fn jit_stats(&self) -> Result<JitQueryStats> {
Ok(lock_safe!(self.jit_stats, "query evaluator jit stats lock")?.clone())
}
pub fn set_jit_enabled(&mut self, enabled: bool) {
self.jit_enabled = enabled;
}
pub fn set_jit_threshold(&mut self, threshold: u64) {
self.jit_threshold = threshold;
}
pub fn clear_jit_cache(&mut self) -> Result<()> {
let mut cache = lock_safe!(
self.compiled_expressions,
"query evaluator compiled expressions lock"
)?;
cache.clear();
Ok(())
}
pub fn compiled_expressions_count(&self) -> Result<usize> {
Ok(lock_safe!(
self.compiled_expressions,
"query evaluator compiled expressions lock"
)?
.len())
}
fn add_builtin_functions(&mut self) {
self.add_function("abs".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args[0].abs()
}
});
self.add_function("sqrt".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args[0].sqrt()
}
});
self.add_function("log".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args[0].ln()
}
});
self.add_function("log10".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args[0].log10()
}
});
self.add_function("exp".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args[0].exp()
}
});
self.add_function("sin".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args[0].sin()
}
});
self.add_function("cos".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args[0].cos()
}
});
self.add_function("tan".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args[0].tan()
}
});
self.add_function("min".to_string(), |args| {
args.iter().fold(f64::INFINITY, |a, &b| a.min(b))
});
self.add_function("max".to_string(), |args| {
args.iter().fold(f64::NEG_INFINITY, |a, &b| a.max(b))
});
self.add_function("sum".to_string(), |args| args.iter().sum());
self.add_function("mean".to_string(), |args| {
if args.is_empty() {
0.0
} else {
args.iter().sum::<f64>() / args.len() as f64
}
});
}
}
pub struct Evaluator<'a> {
dataframe: &'a DataFrame,
context: &'a QueryContext,
column_cache: std::cell::RefCell<HashMap<String, Rc<ValueVec>>>,
enable_short_circuit: bool,
enable_constant_folding: bool,
}
pub struct JitEvaluator<'a> {
dataframe: &'a DataFrame,
context: &'a QueryContext,
}
pub struct OptimizedEvaluator<'a> {
dataframe: &'a DataFrame,
context: &'a QueryContext,
}
impl<'a> Evaluator<'a> {
pub fn new(dataframe: &'a DataFrame, context: &'a QueryContext) -> Self {
Self {
dataframe,
context,
column_cache: std::cell::RefCell::new(HashMap::new()),
enable_short_circuit: true,
enable_constant_folding: true,
}
}
pub fn with_optimizations(
dataframe: &'a DataFrame,
context: &'a QueryContext,
short_circuit: bool,
constant_folding: bool,
) -> Self {
Self {
dataframe,
context,
column_cache: std::cell::RefCell::new(HashMap::new()),
enable_short_circuit: short_circuit,
enable_constant_folding: constant_folding,
}
}
pub fn evaluate_query(&self, expr: &Expr) -> Result<Vec<bool>> {
let prepared = self.prepare(expr)?;
self.evaluate_prepared_query(&prepared)
}
fn evaluate_prepared_query(&self, expr: &Expr) -> Result<Vec<bool>> {
let start = Instant::now();
let row_count = self.dataframe.row_count();
let mut result = Vec::with_capacity(row_count);
for row_idx in 0..row_count {
let value = self.evaluate_expression_for_row(expr, row_idx)?;
result.push(ops::value_to_bool(&value)?);
}
if let Ok(mut stats) = lock_safe!(
self.context.jit_stats,
"query evaluator context jit stats lock"
) {
stats.record_row_by_row_execution(start.elapsed().as_nanos() as u64);
}
Ok(result)
}
pub fn evaluate_query_with_jit(&self, expr: &Expr) -> Result<Vec<bool>> {
let prepared = self.prepared_expression(expr)?;
if self.context.jit_enabled {
let start = Instant::now();
let mask =
VectorizedEvaluator::new(self.dataframe, self.context).evaluate_mask(&prepared)?;
if let Ok(mut stats) = lock_safe!(
self.context.jit_stats,
"query evaluator context jit stats lock"
) {
stats.record_vectorized_execution(start.elapsed().as_nanos() as u64);
}
Ok(mask)
} else {
self.evaluate_prepared_query(&prepared)
}
}
fn prepare(&self, expr: &Expr) -> Result<Expr> {
if self.enable_constant_folding {
self.optimize_expression(expr)
} else {
Ok(expr.clone())
}
}
fn prepared_expression(&self, expr: &Expr) -> Result<Expr> {
let signature = self.expression_signature(expr);
let cache_prepared = {
let mut cache = lock_safe!(
self.context.compiled_expressions,
"query evaluator context compiled expressions lock"
)?;
match cache.get_mut(&signature) {
Some(entry) => {
entry.execution_count += 1;
entry.last_execution = std::time::SystemTime::now();
if let Some(prepared) = &entry.prepared {
return Ok(prepared.clone());
}
entry.execution_count >= self.context.jit_threshold
}
None => {
cache.insert(
signature.clone(),
CompiledExpression {
signature: signature.clone(),
prepared: None,
execution_count: 1,
last_execution: std::time::SystemTime::now(),
},
);
1 >= self.context.jit_threshold
}
}
};
let start = Instant::now();
let prepared = self.prepare(expr)?;
if let Ok(mut stats) = lock_safe!(
self.context.jit_stats,
"query evaluator context jit stats lock"
) {
stats.record_compilation(start.elapsed().as_nanos() as u64);
}
if cache_prepared {
let mut cache = lock_safe!(
self.context.compiled_expressions,
"query evaluator context compiled expressions lock"
)?;
if let Some(entry) = cache.get_mut(&signature) {
entry.prepared = Some(prepared.clone());
}
}
Ok(prepared)
}
fn expression_signature(&self, expr: &Expr) -> String {
format!("{:?}", expr) }
fn optimize_expression(&self, expr: &Expr) -> Result<Expr> {
match expr {
Expr::Binary { left, op, right } => {
let optimized_left = self.optimize_expression(left)?;
let optimized_right = self.optimize_expression(right)?;
if let (Expr::Literal(l), Expr::Literal(r)) = (&optimized_left, &optimized_right) {
let result = self.apply_binary_operation(l, op, r)?;
return Ok(Expr::Literal(result));
}
match (&optimized_left, op, &optimized_right) {
(expr, BinaryOp::And, Expr::Literal(LiteralValue::Boolean(true))) => {
Ok(expr.clone())
}
(Expr::Literal(LiteralValue::Boolean(true)), BinaryOp::And, expr) => {
Ok(expr.clone())
}
(expr, BinaryOp::Or, Expr::Literal(LiteralValue::Boolean(false))) => {
Ok(expr.clone())
}
(Expr::Literal(LiteralValue::Boolean(false)), BinaryOp::Or, expr) => {
Ok(expr.clone())
}
(expr, BinaryOp::Add, Expr::Literal(LiteralValue::Number(n))) if *n == 0.0 => {
Ok(expr.clone())
}
(Expr::Literal(LiteralValue::Number(n)), BinaryOp::Add, expr) if *n == 0.0 => {
Ok(expr.clone())
}
(expr, BinaryOp::Multiply, Expr::Literal(LiteralValue::Number(n)))
if *n == 1.0 =>
{
Ok(expr.clone())
}
(Expr::Literal(LiteralValue::Number(n)), BinaryOp::Multiply, expr)
if *n == 1.0 =>
{
Ok(expr.clone())
}
_ => Ok(Expr::Binary {
left: Box::new(optimized_left),
op: op.clone(),
right: Box::new(optimized_right),
}),
}
}
Expr::Unary { op, operand } => {
let optimized_operand = self.optimize_expression(operand)?;
if let Expr::Literal(val) = &optimized_operand {
let result = self.apply_unary_operation(op, val)?;
return Ok(Expr::Literal(result));
}
if let (
UnaryOp::Not,
Expr::Unary {
op: UnaryOp::Not,
operand,
},
) = (op, &optimized_operand)
{
return Ok((**operand).clone());
}
Ok(Expr::Unary {
op: op.clone(),
operand: Box::new(optimized_operand),
})
}
Expr::Function { name, args } => {
let optimized_args: Result<Vec<Expr>> = args
.iter()
.map(|arg| self.optimize_expression(arg))
.collect();
Ok(Expr::Function {
name: name.clone(),
args: optimized_args?,
})
}
_ => Ok(expr.clone()),
}
}
pub fn evaluate_expression_for_row(&self, expr: &Expr, row_idx: usize) -> Result<LiteralValue> {
match expr {
Expr::Literal(value) => Ok(value.clone()),
Expr::Variable(name) => match self.context.variables.get(name) {
Some(value) => Ok(value.clone()),
None => Err(Error::InvalidValue(format!(
"Undefined query variable '@{}'",
name
))),
},
Expr::Column(name) => {
{
let cache = self.column_cache.borrow();
if let Some(cached) = cache.get(name) {
return cached.cell_at(row_idx);
}
}
if !self.dataframe.contains_column(name) {
if let Some(value) = self.context.variables.get(name) {
return Ok(value.clone());
}
return Err(Error::ColumnNotFound(name.clone()));
}
let loaded = Rc::new(load_column_values(self.dataframe, name)?);
let value = loaded.cell_at(row_idx);
self.column_cache
.borrow_mut()
.insert(name.clone(), Rc::clone(&loaded));
value
}
Expr::Binary { left, op, right } => {
if self.enable_short_circuit {
match op {
BinaryOp::And => {
let left_val = self.evaluate_expression_for_row(left, row_idx)?;
if let LiteralValue::Boolean(false) = left_val {
return Ok(LiteralValue::Boolean(false)); }
let right_val = self.evaluate_expression_for_row(right, row_idx)?;
self.apply_binary_operation(&left_val, op, &right_val)
}
BinaryOp::Or => {
let left_val = self.evaluate_expression_for_row(left, row_idx)?;
if let LiteralValue::Boolean(true) = left_val {
return Ok(LiteralValue::Boolean(true)); }
let right_val = self.evaluate_expression_for_row(right, row_idx)?;
self.apply_binary_operation(&left_val, op, &right_val)
}
_ => {
let left_val = self.evaluate_expression_for_row(left, row_idx)?;
let right_val = self.evaluate_expression_for_row(right, row_idx)?;
self.apply_binary_operation(&left_val, op, &right_val)
}
}
} else {
let left_val = self.evaluate_expression_for_row(left, row_idx)?;
let right_val = self.evaluate_expression_for_row(right, row_idx)?;
self.apply_binary_operation(&left_val, op, &right_val)
}
}
Expr::Unary { op, operand } => {
let operand_val = self.evaluate_expression_for_row(operand, row_idx)?;
self.apply_unary_operation(op, &operand_val)
}
Expr::Function { name, args } => {
let arg_values: Result<Vec<f64>> = args
.iter()
.map(|arg| {
let val = self.evaluate_expression_for_row(arg, row_idx)?;
match val {
LiteralValue::Number(n) => Ok(n),
LiteralValue::String(s) => ops::text_to_number(&s)
.ok_or_else(|| ops::non_numeric_text_error(&s)),
LiteralValue::Boolean(_) => Err(Error::InvalidValue(
"Function arguments must be numeric".to_string(),
)),
}
})
.collect();
let arg_values = arg_values?;
if let Some(func) = self.context.functions.get(name) {
let result = func(&arg_values);
Ok(LiteralValue::Number(result))
} else {
Err(Error::InvalidValue(format!("Unknown function: {}", name)))
}
}
}
}
fn apply_binary_operation(
&self,
left: &LiteralValue,
op: &BinaryOp,
right: &LiteralValue,
) -> Result<LiteralValue> {
ops::apply_binary(left, op, right)
}
fn apply_unary_operation(&self, op: &UnaryOp, operand: &LiteralValue) -> Result<LiteralValue> {
ops::apply_unary(op, operand)
}
}
impl<'a> JitEvaluator<'a> {
pub fn new(dataframe: &'a DataFrame, context: &'a QueryContext) -> Self {
Self { dataframe, context }
}
pub fn evaluate_query_jit(&self, expr: &Expr) -> Result<Vec<bool>> {
Evaluator::new(self.dataframe, self.context).evaluate_query_with_jit(expr)
}
}
impl<'a> OptimizedEvaluator<'a> {
pub fn new(dataframe: &'a DataFrame, context: &'a QueryContext) -> Self {
Self { dataframe, context }
}
pub fn evaluate_query_vectorized(&self, expr: &Expr) -> Result<Vec<bool>> {
let start = Instant::now();
let mask = VectorizedEvaluator::new(self.dataframe, self.context).evaluate_mask(expr)?;
if let Ok(mut stats) = lock_safe!(
self.context.jit_stats,
"query evaluator context jit stats lock"
) {
stats.record_vectorized_execution(start.elapsed().as_nanos() as u64);
}
Ok(mask)
}
}
impl Clone for QueryContext {
fn clone(&self) -> Self {
let mut new_context = Self {
variables: self.variables.clone(),
functions: HashMap::new(),
compiled_expressions: Arc::clone(&self.compiled_expressions),
jit_stats: Arc::clone(&self.jit_stats),
jit_threshold: self.jit_threshold,
jit_enabled: self.jit_enabled,
};
new_context.add_builtin_functions();
new_context
}
}