Skip to main content

xlog_runtime/executor/
expression.rs

1//! Expression evaluation methods for the Executor.
2//!
3//! Production GPU-accelerated filter, predicate mask, arithmetic expression,
4//! and mask operation methods.
5
6use cudarc::driver::LaunchConfig;
7use xlog_core::{Result, ScalarType, Schema, XlogError};
8use xlog_cuda::memory::TrackedCudaSlice;
9use xlog_cuda::provider::{arith_kernels, filter_kernels, ARITH_MODULE, FILTER_MODULE};
10use xlog_cuda::{CudaBuffer, LaunchAsync};
11use xlog_ir::{CompareOp, ConstValue, Expr, ProjectExpr};
12
13use super::Executor;
14
15#[derive(Clone, Copy)]
16enum ArithmeticBinaryOperation {
17    Add,
18    Subtract,
19    Multiply,
20    Divide,
21    Modulo,
22    Minimum,
23    Maximum,
24    Power,
25}
26
27#[derive(Clone, Copy)]
28enum MaskBinaryOperation {
29    And,
30    Or,
31}
32
33enum ExpressionTask<'a> {
34    Arithmetic(&'a Expr),
35    Predicate(&'a Expr),
36    FinishArithmeticBinary(ArithmeticBinaryOperation),
37    FinishAbsoluteValue,
38    FinishCast(ScalarType),
39    FinishComparison {
40        op: CompareOp,
41        use_float: bool,
42    },
43    ContinueMaskFold {
44        expressions: &'a [Expr],
45        next_index: usize,
46        operation: MaskBinaryOperation,
47    },
48    FinishMaskFold {
49        expressions: &'a [Expr],
50        next_index: usize,
51        operation: MaskBinaryOperation,
52    },
53    FinishMaskNot,
54    PrepareConditional {
55        then_expr: &'a Expr,
56        else_expr: &'a Expr,
57    },
58    FinishConditional,
59}
60
61enum ExpressionValue {
62    Arithmetic(CudaBuffer),
63    Predicate(TrackedCudaSlice<u8>),
64    SelectionMask(CudaBuffer),
65}
66
67impl Executor {
68    /// Check if an expression may produce a floating-point result.
69    pub(crate) fn expr_may_be_float(expr: &Expr, schema: &Schema) -> bool {
70        let mut pending = vec![expr];
71        while let Some(current) = pending.pop() {
72            match current {
73                Expr::Column(col_idx)
74                    if matches!(
75                        schema.column_type(*col_idx),
76                        Some(ScalarType::F32 | ScalarType::F64)
77                    ) =>
78                {
79                    return true;
80                }
81                Expr::Const(ConstValue::F32(_) | ConstValue::F64(_))
82                | Expr::Cast(_, ScalarType::F32 | ScalarType::F64) => return true,
83                Expr::Add(left, right)
84                | Expr::Sub(left, right)
85                | Expr::Mul(left, right)
86                | Expr::Div(left, right)
87                | Expr::Mod(left, right)
88                | Expr::Min(left, right)
89                | Expr::Max(left, right)
90                | Expr::Pow(left, right) => {
91                    pending.push(right);
92                    pending.push(left);
93                }
94                Expr::Abs(inner) | Expr::Cast(inner, _) => pending.push(inner),
95                _ => {}
96            }
97        }
98        false
99    }
100
101    /// Execute a Filter node using GPU predicate evaluation.
102    pub fn execute_filter(&self, input: &CudaBuffer, predicate: &Expr) -> Result<CudaBuffer> {
103        if input.is_empty() {
104            return self.create_empty_buffer(input.schema().clone());
105        }
106
107        let mask = self.eval_predicate_mask_gpu(predicate, input)?;
108        self.provider.filter_by_device_mask(input, &mask)
109    }
110
111    pub(crate) fn eval_predicate_mask_gpu(
112        &self,
113        expr: &Expr,
114        input: &CudaBuffer,
115    ) -> Result<TrackedCudaSlice<u8>> {
116        match self.evaluate_expression(expr, input, true)? {
117            ExpressionValue::Predicate(mask) => Ok(mask),
118            _ => Err(Self::expression_state_error(
119                "predicate evaluation produced an arithmetic value",
120            )),
121        }
122    }
123
124    fn compare_buffers_mask(
125        &self,
126        left: &CudaBuffer,
127        right: &CudaBuffer,
128        op: CompareOp,
129    ) -> Result<TrackedCudaSlice<u8>> {
130        if left.arity() != 1 || right.arity() != 1 {
131            return Err(XlogError::Execution(
132                "Compare requires single-column buffers".into(),
133            ));
134        }
135        if left.num_rows() != right.num_rows() {
136            return Err(XlogError::Execution(
137                "Compare requires matching row counts".into(),
138            ));
139        }
140        if left.num_rows() > u32::MAX as u64 {
141            return Err(XlogError::Execution(format!(
142                "Compare supports at most {} rows, got {}",
143                u32::MAX,
144                left.num_rows()
145            )));
146        }
147        if left.is_empty() {
148            return self.provider.memory().alloc::<u8>(0).map_err(|e| {
149                XlogError::execution_ctx("compare_buffers_mask", "allocate empty mask", &e)
150            });
151        }
152
153        let left_type = left
154            .schema()
155            .column_type(0)
156            .ok_or_else(|| XlogError::Execution("Missing left column type".into()))?;
157        let right_type = right
158            .schema()
159            .column_type(0)
160            .ok_or_else(|| XlogError::Execution("Missing right column type".into()))?;
161
162        if left_type != right_type {
163            return Err(XlogError::Execution(
164                "Compare requires matching column types".into(),
165            ));
166        }
167
168        let kernel = match left_type {
169            ScalarType::U32 | ScalarType::Symbol => filter_kernels::FILTER_COMPARE_U32_COL,
170            ScalarType::U64 => filter_kernels::FILTER_COMPARE_U64_COL,
171            ScalarType::I32 => filter_kernels::FILTER_COMPARE_I32_COL,
172            ScalarType::I64 => filter_kernels::FILTER_COMPARE_I64_COL,
173            ScalarType::F32 => filter_kernels::FILTER_COMPARE_F32_COL,
174            ScalarType::F64 => filter_kernels::FILTER_COMPARE_F64_COL,
175            ScalarType::Bool => filter_kernels::FILTER_COMPARE_U8_COL,
176        };
177
178        let left_col = left
179            .column(0)
180            .ok_or_else(|| XlogError::Execution("Missing left column".into()))?;
181        let right_col = right
182            .column(0)
183            .ok_or_else(|| XlogError::Execution("Missing right column".into()))?;
184
185        let num_rows = left.num_rows() as u32;
186        let mut d_mask = self.provider.memory().alloc::<u8>(num_rows as usize)?;
187
188        let func = self
189            .provider
190            .device()
191            .inner()
192            .get_func(FILTER_MODULE, kernel)
193            .ok_or_else(|| XlogError::Execution("filter compare kernel not found".into()))?;
194        let config = LaunchConfig::for_num_elems(num_rows);
195
196        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
197        unsafe {
198            func.clone().launch(
199                config,
200                (left_col, right_col, num_rows, op as u8, &mut d_mask),
201            )
202        }
203        .map_err(|e| XlogError::execution_ctx("compare_buffers_mask", "filter compare", &e))?;
204
205        Ok(d_mask)
206    }
207
208    fn mask_and(
209        &self,
210        left: &TrackedCudaSlice<u8>,
211        right: &TrackedCudaSlice<u8>,
212        n: u32,
213    ) -> Result<TrackedCudaSlice<u8>> {
214        let mut out = self.provider.memory().alloc::<u8>(n as usize)?;
215        if n == 0 {
216            return Ok(out);
217        }
218
219        let func = self
220            .provider
221            .device()
222            .inner()
223            .get_func(FILTER_MODULE, filter_kernels::MASK_AND)
224            .ok_or_else(|| XlogError::Execution("mask_and kernel not found".into()))?;
225        let config = LaunchConfig::for_num_elems(n);
226
227        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
228        unsafe { func.clone().launch(config, (left, right, &mut out, n)) }
229            .map_err(|e| XlogError::execution_ctx("mask_and", "launch kernel", &e))?;
230
231        Ok(out)
232    }
233
234    fn mask_or(
235        &self,
236        left: &TrackedCudaSlice<u8>,
237        right: &TrackedCudaSlice<u8>,
238        n: u32,
239    ) -> Result<TrackedCudaSlice<u8>> {
240        let mut out = self.provider.memory().alloc::<u8>(n as usize)?;
241        if n == 0 {
242            return Ok(out);
243        }
244
245        let func = self
246            .provider
247            .device()
248            .inner()
249            .get_func(FILTER_MODULE, filter_kernels::MASK_OR)
250            .ok_or_else(|| XlogError::Execution("mask_or kernel not found".into()))?;
251        let config = LaunchConfig::for_num_elems(n);
252
253        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
254        unsafe { func.clone().launch(config, (left, right, &mut out, n)) }
255            .map_err(|e| XlogError::execution_ctx("mask_or", "launch kernel", &e))?;
256
257        Ok(out)
258    }
259
260    fn mask_not(&self, input: &TrackedCudaSlice<u8>, n: u32) -> Result<TrackedCudaSlice<u8>> {
261        let mut out = self.provider.memory().alloc::<u8>(n as usize)?;
262        if n == 0 {
263            return Ok(out);
264        }
265
266        let func = self
267            .provider
268            .device()
269            .inner()
270            .get_func(FILTER_MODULE, filter_kernels::MASK_NOT)
271            .ok_or_else(|| XlogError::Execution("mask_not kernel not found".into()))?;
272        let config = LaunchConfig::for_num_elems(n);
273
274        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
275        unsafe { func.clone().launch(config, (input, &mut out, n)) }
276            .map_err(|e| XlogError::execution_ctx("mask_not", "launch kernel", &e))?;
277
278        Ok(out)
279    }
280
281    fn mask_filled(&self, n: u32, value: u8) -> Result<TrackedCudaSlice<u8>> {
282        let mut out = self.provider.memory().alloc::<u8>(n as usize)?;
283        if n == 0 {
284            return Ok(out);
285        }
286
287        if value == 0 {
288            self.provider
289                .device()
290                .inner()
291                .memset_zeros(&mut out)
292                .map_err(|e| XlogError::execution_ctx("mask_filled", "mask memset", &e))?;
293            return Ok(out);
294        }
295
296        let func = self
297            .provider
298            .device()
299            .inner()
300            .get_func(ARITH_MODULE, arith_kernels::ARITH_FILL_CONST_U8)
301            .ok_or_else(|| XlogError::Execution("arith fill kernel not found".into()))?;
302        let config = LaunchConfig::for_num_elems(n);
303
304        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
305        unsafe { func.clone().launch(config, (value, n, &mut out)) }
306            .map_err(|e| XlogError::execution_ctx("mask_filled", "mask fill", &e))?;
307
308        Ok(out)
309    }
310
311    pub(crate) fn wrap_single_column(
312        &self,
313        buffer: &CudaBuffer,
314        col_idx: usize,
315    ) -> Result<CudaBuffer> {
316        let col_type = buffer
317            .schema()
318            .column_type(col_idx)
319            .ok_or_else(|| XlogError::Execution(format!("Column {} not found", col_idx)))?;
320        let schema = Schema::new(vec![("expr".to_string(), col_type)]);
321
322        if buffer.is_empty() {
323            return self.create_empty_buffer(schema);
324        }
325
326        let num_rows = buffer.num_rows();
327        let bytes = (num_rows as usize)
328            .checked_mul(col_type.size_bytes())
329            .ok_or_else(|| XlogError::Execution("Column size overflow".into()))?;
330
331        let src_col = buffer
332            .column(col_idx)
333            .ok_or_else(|| XlogError::Execution(format!("Column {} not found", col_idx)))?;
334        let mut dst_col = self.provider.memory().alloc::<u8>(bytes)?;
335        if bytes > 0 {
336            self.provider
337                .device()
338                .inner()
339                .dtod_copy(src_col, &mut dst_col)
340                .map_err(|e| XlogError::execution_ctx("wrap_single_column", "copy column", &e))?;
341        }
342
343        let d_num_rows = self.clone_device_row_count(buffer)?;
344        self.provider.device().synchronize()?;
345        Ok(CudaBuffer::from_columns(
346            vec![dst_col.into()],
347            num_rows,
348            d_num_rows,
349            schema,
350        ))
351    }
352
353    /// Evaluate an arithmetic expression on a buffer, producing a single-column result.
354    ///
355    /// The explicit task stack preserves source-order evaluation without consuming
356    /// one native stack frame per nested expression.
357    pub(crate) fn evaluate_arith_expr(
358        &self,
359        expr: &Expr,
360        input: &CudaBuffer,
361    ) -> Result<CudaBuffer> {
362        match self.evaluate_expression(expr, input, false)? {
363            ExpressionValue::Arithmetic(buffer) => Ok(buffer),
364            _ => Err(Self::expression_state_error(
365                "arithmetic evaluation produced a predicate value",
366            )),
367        }
368    }
369
370    fn evaluate_expression(
371        &self,
372        expression: &Expr,
373        input: &CudaBuffer,
374        predicate: bool,
375    ) -> Result<ExpressionValue> {
376        let mut tasks = vec![if predicate {
377            ExpressionTask::Predicate(expression)
378        } else {
379            ExpressionTask::Arithmetic(expression)
380        }];
381        let mut values = Vec::new();
382
383        while let Some(task) = tasks.pop() {
384            match task {
385                ExpressionTask::Arithmetic(expression) => match expression {
386                    Expr::Column(index) => {
387                        values.push(ExpressionValue::Arithmetic(
388                            self.wrap_single_column(input, *index)?,
389                        ));
390                    }
391                    Expr::Const(value) => {
392                        let (bytes, column_type) = self.const_to_bytes_and_type(value);
393                        values.push(ExpressionValue::Arithmetic(
394                            self.provider.create_constant_column_with_device_count(
395                                &bytes,
396                                column_type,
397                                input.num_rows(),
398                                input.num_rows_device(),
399                            )?,
400                        ));
401                    }
402                    Expr::Add(left, right) => Self::schedule_arithmetic_binary(
403                        &mut tasks,
404                        left,
405                        right,
406                        ArithmeticBinaryOperation::Add,
407                    ),
408                    Expr::Sub(left, right) => Self::schedule_arithmetic_binary(
409                        &mut tasks,
410                        left,
411                        right,
412                        ArithmeticBinaryOperation::Subtract,
413                    ),
414                    Expr::Mul(left, right) => Self::schedule_arithmetic_binary(
415                        &mut tasks,
416                        left,
417                        right,
418                        ArithmeticBinaryOperation::Multiply,
419                    ),
420                    Expr::Div(left, right) => Self::schedule_arithmetic_binary(
421                        &mut tasks,
422                        left,
423                        right,
424                        ArithmeticBinaryOperation::Divide,
425                    ),
426                    Expr::Mod(left, right) => Self::schedule_arithmetic_binary(
427                        &mut tasks,
428                        left,
429                        right,
430                        ArithmeticBinaryOperation::Modulo,
431                    ),
432                    Expr::Min(left, right) => Self::schedule_arithmetic_binary(
433                        &mut tasks,
434                        left,
435                        right,
436                        ArithmeticBinaryOperation::Minimum,
437                    ),
438                    Expr::Max(left, right) => Self::schedule_arithmetic_binary(
439                        &mut tasks,
440                        left,
441                        right,
442                        ArithmeticBinaryOperation::Maximum,
443                    ),
444                    Expr::Pow(left, right) => Self::schedule_arithmetic_binary(
445                        &mut tasks,
446                        left,
447                        right,
448                        ArithmeticBinaryOperation::Power,
449                    ),
450                    Expr::Abs(inner) => {
451                        tasks.push(ExpressionTask::FinishAbsoluteValue);
452                        tasks.push(ExpressionTask::Arithmetic(inner));
453                    }
454                    Expr::Cast(inner, target) => {
455                        tasks.push(ExpressionTask::FinishCast(*target));
456                        tasks.push(ExpressionTask::Arithmetic(inner));
457                    }
458                    Expr::Conditional {
459                        condition,
460                        then_expr,
461                        else_expr,
462                    } => {
463                        tasks.push(ExpressionTask::PrepareConditional {
464                            then_expr,
465                            else_expr,
466                        });
467                        tasks.push(ExpressionTask::Predicate(condition));
468                    }
469                    Expr::Compare { .. } | Expr::And(_) | Expr::Or(_) | Expr::Not(_) => {
470                        return Err(XlogError::Execution(format!(
471                            "Unsupported expression in arithmetic evaluation: {:?}",
472                            expression
473                        )));
474                    }
475                },
476                ExpressionTask::Predicate(expression) => {
477                    if input.num_rows() > u32::MAX as u64 {
478                        return Err(XlogError::Execution(format!(
479                            "Predicate evaluation supports at most {} rows, got {}",
480                            u32::MAX,
481                            input.num_rows()
482                        )));
483                    }
484                    let row_count = input.num_rows() as u32;
485                    match expression {
486                        Expr::Column(column_index) => {
487                            let column_type =
488                                input.schema().column_type(*column_index).ok_or_else(|| {
489                                    XlogError::Execution(format!(
490                                        "Column {} not found",
491                                        column_index
492                                    ))
493                                })?;
494                            let mask = if column_type == ScalarType::Bool {
495                                let column = self.wrap_single_column(input, *column_index)?;
496                                let zero = self.provider.create_constant_column_with_device_count(
497                                    &[0u8],
498                                    ScalarType::Bool,
499                                    input.num_rows(),
500                                    input.num_rows_device(),
501                                )?;
502                                self.compare_buffers_mask(&column, &zero, CompareOp::Ne)?
503                            } else {
504                                self.mask_filled(row_count, 1)?
505                            };
506                            values.push(ExpressionValue::Predicate(mask));
507                        }
508                        Expr::Const(ConstValue::Bool(value)) => {
509                            values.push(ExpressionValue::Predicate(
510                                self.mask_filled(row_count, u8::from(*value))?,
511                            ))
512                        }
513                        Expr::Const(_) => {
514                            values.push(ExpressionValue::Predicate(self.mask_filled(row_count, 1)?))
515                        }
516                        Expr::Compare { left, op, right } => {
517                            let use_float = Self::expr_may_be_float(left, input.schema())
518                                || Self::expr_may_be_float(right, input.schema());
519                            tasks.push(ExpressionTask::FinishComparison { op: *op, use_float });
520                            tasks.push(ExpressionTask::Arithmetic(right));
521                            tasks.push(ExpressionTask::Arithmetic(left));
522                        }
523                        Expr::And(expressions) => Self::schedule_mask_fold(
524                            &mut tasks,
525                            &mut values,
526                            expressions,
527                            row_count,
528                            MaskBinaryOperation::And,
529                            self,
530                        )?,
531                        Expr::Or(expressions) => Self::schedule_mask_fold(
532                            &mut tasks,
533                            &mut values,
534                            expressions,
535                            row_count,
536                            MaskBinaryOperation::Or,
537                            self,
538                        )?,
539                        Expr::Not(inner) => {
540                            tasks.push(ExpressionTask::FinishMaskNot);
541                            tasks.push(ExpressionTask::Predicate(inner));
542                        }
543                        Expr::Add(_, _)
544                        | Expr::Sub(_, _)
545                        | Expr::Mul(_, _)
546                        | Expr::Div(_, _)
547                        | Expr::Mod(_, _)
548                        | Expr::Abs(_)
549                        | Expr::Min(_, _)
550                        | Expr::Max(_, _)
551                        | Expr::Pow(_, _)
552                        | Expr::Cast(_, _)
553                        | Expr::Conditional { .. } => {
554                            return Err(XlogError::Execution(
555                                "Arithmetic expression cannot be evaluated as boolean predicate"
556                                    .into(),
557                            ));
558                        }
559                    }
560                }
561                ExpressionTask::FinishArithmeticBinary(operation) => {
562                    let right = Self::pop_arithmetic_value(&mut values)?;
563                    let left = Self::pop_arithmetic_value(&mut values)?;
564                    let result = match operation {
565                        ArithmeticBinaryOperation::Add => self.provider.add_columns(&left, &right),
566                        ArithmeticBinaryOperation::Subtract => {
567                            self.provider.sub_columns(&left, &right)
568                        }
569                        ArithmeticBinaryOperation::Multiply => {
570                            self.provider.mul_columns(&left, &right)
571                        }
572                        ArithmeticBinaryOperation::Divide => {
573                            self.provider.div_columns(&left, &right)
574                        }
575                        ArithmeticBinaryOperation::Modulo => {
576                            self.provider.mod_columns(&left, &right)
577                        }
578                        ArithmeticBinaryOperation::Minimum => {
579                            self.provider.min_columns(&left, &right)
580                        }
581                        ArithmeticBinaryOperation::Maximum => {
582                            self.provider.max_columns(&left, &right)
583                        }
584                        ArithmeticBinaryOperation::Power => {
585                            self.provider.pow_columns(&left, &right)
586                        }
587                    }?;
588                    values.push(ExpressionValue::Arithmetic(result));
589                }
590                ExpressionTask::FinishAbsoluteValue => {
591                    let value = Self::pop_arithmetic_value(&mut values)?;
592                    values.push(ExpressionValue::Arithmetic(
593                        self.provider.abs_column(&value)?,
594                    ));
595                }
596                ExpressionTask::FinishCast(target) => {
597                    let value = Self::pop_arithmetic_value(&mut values)?;
598                    values.push(ExpressionValue::Arithmetic(
599                        self.provider.cast_column(&value, target)?,
600                    ));
601                }
602                ExpressionTask::FinishComparison { op, use_float } => {
603                    let mut right = Self::pop_arithmetic_value(&mut values)?;
604                    let mut left = Self::pop_arithmetic_value(&mut values)?;
605                    if use_float {
606                        left = self.provider.cast_column(&left, ScalarType::F64)?;
607                        right = self.provider.cast_column(&right, ScalarType::F64)?;
608                    }
609                    values.push(ExpressionValue::Predicate(
610                        self.compare_buffers_mask(&left, &right, op)?,
611                    ));
612                }
613                ExpressionTask::ContinueMaskFold {
614                    expressions,
615                    next_index,
616                    operation,
617                } => {
618                    if next_index < expressions.len() {
619                        tasks.push(ExpressionTask::FinishMaskFold {
620                            expressions,
621                            next_index: next_index + 1,
622                            operation,
623                        });
624                        tasks.push(ExpressionTask::Predicate(&expressions[next_index]));
625                    }
626                }
627                ExpressionTask::FinishMaskFold {
628                    expressions,
629                    next_index,
630                    operation,
631                } => {
632                    let right = Self::pop_predicate_value(&mut values)?;
633                    let left = Self::pop_predicate_value(&mut values)?;
634                    let row_count = input.num_rows() as u32;
635                    let combined = match operation {
636                        MaskBinaryOperation::And => self.mask_and(&left, &right, row_count),
637                        MaskBinaryOperation::Or => self.mask_or(&left, &right, row_count),
638                    }?;
639                    values.push(ExpressionValue::Predicate(combined));
640                    tasks.push(ExpressionTask::ContinueMaskFold {
641                        expressions,
642                        next_index,
643                        operation,
644                    });
645                }
646                ExpressionTask::FinishMaskNot => {
647                    let mask = Self::pop_predicate_value(&mut values)?;
648                    values.push(ExpressionValue::Predicate(
649                        self.mask_not(&mask, input.num_rows() as u32)?,
650                    ));
651                }
652                ExpressionTask::PrepareConditional {
653                    then_expr,
654                    else_expr,
655                } => {
656                    let mask = Self::pop_predicate_value(&mut values)?;
657                    let device_row_count = self.clone_device_row_count(input)?;
658                    values.push(ExpressionValue::SelectionMask(CudaBuffer::from_columns(
659                        vec![mask.into()],
660                        input.num_rows(),
661                        device_row_count,
662                        Schema::new(vec![("mask".to_string(), ScalarType::Bool)]),
663                    )));
664                    tasks.push(ExpressionTask::FinishConditional);
665                    tasks.push(ExpressionTask::Arithmetic(else_expr));
666                    tasks.push(ExpressionTask::Arithmetic(then_expr));
667                }
668                ExpressionTask::FinishConditional => {
669                    let else_value = Self::pop_arithmetic_value(&mut values)?;
670                    let then_value = Self::pop_arithmetic_value(&mut values)?;
671                    let mask = match values.pop() {
672                        Some(ExpressionValue::SelectionMask(mask)) => mask,
673                        _ => {
674                            return Err(Self::expression_state_error(
675                                "conditional evaluation is missing its selection mask",
676                            ));
677                        }
678                    };
679                    values.push(ExpressionValue::Arithmetic(self.provider.select_columns(
680                        &mask,
681                        &then_value,
682                        &else_value,
683                    )?));
684                }
685            }
686        }
687
688        if values.len() != 1 {
689            return Err(Self::expression_state_error(
690                "expression evaluation did not produce exactly one value",
691            ));
692        }
693        values
694            .pop()
695            .ok_or_else(|| Self::expression_state_error("expression evaluation produced no value"))
696    }
697
698    fn schedule_arithmetic_binary<'a>(
699        tasks: &mut Vec<ExpressionTask<'a>>,
700        left: &'a Expr,
701        right: &'a Expr,
702        operation: ArithmeticBinaryOperation,
703    ) {
704        tasks.push(ExpressionTask::FinishArithmeticBinary(operation));
705        tasks.push(ExpressionTask::Arithmetic(right));
706        tasks.push(ExpressionTask::Arithmetic(left));
707    }
708
709    fn schedule_mask_fold<'a>(
710        tasks: &mut Vec<ExpressionTask<'a>>,
711        values: &mut Vec<ExpressionValue>,
712        expressions: &'a [Expr],
713        row_count: u32,
714        operation: MaskBinaryOperation,
715        executor: &Self,
716    ) -> Result<()> {
717        if expressions.is_empty() {
718            let identity = match operation {
719                MaskBinaryOperation::And => 1,
720                MaskBinaryOperation::Or => 0,
721            };
722            values.push(ExpressionValue::Predicate(
723                executor.mask_filled(row_count, identity)?,
724            ));
725        } else {
726            tasks.push(ExpressionTask::ContinueMaskFold {
727                expressions,
728                next_index: 1,
729                operation,
730            });
731            tasks.push(ExpressionTask::Predicate(&expressions[0]));
732        }
733        Ok(())
734    }
735
736    fn pop_arithmetic_value(values: &mut Vec<ExpressionValue>) -> Result<CudaBuffer> {
737        match values.pop() {
738            Some(ExpressionValue::Arithmetic(value)) => Ok(value),
739            _ => Err(Self::expression_state_error(
740                "arithmetic operation is missing an operand",
741            )),
742        }
743    }
744
745    fn pop_predicate_value(values: &mut Vec<ExpressionValue>) -> Result<TrackedCudaSlice<u8>> {
746        match values.pop() {
747            Some(ExpressionValue::Predicate(value)) => Ok(value),
748            _ => Err(Self::expression_state_error(
749                "predicate operation is missing an operand",
750            )),
751        }
752    }
753
754    fn expression_state_error(message: &str) -> XlogError {
755        XlogError::Execution(format!("Internal expression evaluator error: {message}"))
756    }
757
758    /// Convert a ConstValue to raw bytes and ScalarType
759    pub(crate) fn const_to_bytes_and_type(&self, val: &ConstValue) -> (Vec<u8>, ScalarType) {
760        match val {
761            ConstValue::U32(v) => (v.to_le_bytes().to_vec(), ScalarType::U32),
762            ConstValue::U64(v) => (v.to_le_bytes().to_vec(), ScalarType::U64),
763            ConstValue::I32(v) => (v.to_le_bytes().to_vec(), ScalarType::I32),
764            ConstValue::I64(v) => (v.to_le_bytes().to_vec(), ScalarType::I64),
765            ConstValue::F32(v) => (v.to_le_bytes().to_vec(), ScalarType::F32),
766            ConstValue::F64(v) => (v.to_le_bytes().to_vec(), ScalarType::F64),
767            ConstValue::Bool(v) => (vec![if *v { 1u8 } else { 0u8 }], ScalarType::Bool),
768            ConstValue::Symbol(s) => (
769                xlog_core::symbol::intern(s).to_le_bytes().to_vec(),
770                ScalarType::Symbol,
771            ),
772        }
773    }
774
775    /// Execute a Project node
776    ///
777    /// Selects and reorders columns according to the projection list.
778    /// Supports both column pass-through and computed expressions.
779    pub(crate) fn execute_project(
780        &self,
781        input: &CudaBuffer,
782        columns: &[ProjectExpr],
783    ) -> Result<CudaBuffer> {
784        if input.is_empty() {
785            // Build projected schema
786            let projected_schema = self.project_schema(input.schema(), columns)?;
787            return self.create_empty_buffer(projected_schema);
788        }
789
790        if columns.is_empty() {
791            // A zero-column projection preserves row existence. Combining an empty
792            // list of column buffers would manufacture a zero-row relation and make
793            // ground negation treat every matching atom as absent.
794            let projected_schema = self.project_schema(input.schema(), columns)?;
795            let rows = self.provider.device_row_count(input)?;
796            let rows = u32::try_from(rows).map_err(|_| {
797                XlogError::Execution(format!(
798                    "zero-column projection row count {rows} exceeds the GPU range"
799                ))
800            })?;
801            return self
802                .provider
803                .create_zero_arity_buffer(projected_schema, rows);
804        }
805
806        // Build result columns as single-column CudaBuffers
807        let mut result_buffers: Vec<CudaBuffer> = Vec::with_capacity(columns.len());
808        let mut result_types: Vec<ScalarType> = Vec::with_capacity(columns.len());
809
810        for proj_expr in columns {
811            match proj_expr {
812                ProjectExpr::Column(col_idx) => {
813                    // Use extract_column to get a single-column buffer
814                    let col_buffer = self.provider.extract_column(input, *col_idx)?;
815                    let col_type = input
816                        .schema()
817                        .column_type(*col_idx)
818                        .unwrap_or(ScalarType::U64);
819                    result_types.push(col_type);
820                    result_buffers.push(col_buffer);
821                }
822                ProjectExpr::Computed(expr, result_type) => {
823                    // Evaluate the arithmetic expression to get a single-column buffer
824                    let computed_buffer = self.evaluate_arith_expr(expr, input)?;
825                    result_types.push(*result_type);
826                    result_buffers.push(computed_buffer);
827                }
828            }
829        }
830
831        let projected_schema = self.project_schema(input.schema(), columns)?;
832        let mut output = self
833            .provider
834            .combine_columns(result_buffers, result_types)?;
835        output.set_schema(projected_schema);
836        Ok(output)
837    }
838
839    /// Build a projected schema from ProjectExpr list
840    pub(crate) fn project_schema(&self, input: &Schema, columns: &[ProjectExpr]) -> Result<Schema> {
841        let mut projected_columns: Vec<(String, ScalarType)> = Vec::with_capacity(columns.len());
842        let mut projected_sort_labels: Vec<String> = Vec::with_capacity(columns.len());
843        for proj_expr in columns {
844            match proj_expr {
845                ProjectExpr::Column(col_idx) => {
846                    if let Some((name, ty)) = input.columns.get(*col_idx) {
847                        projected_columns.push((name.clone(), *ty));
848                        projected_sort_labels.push(
849                            input
850                                .column_sort_label(*col_idx)
851                                .unwrap_or(name)
852                                .to_string(),
853                        );
854                    } else {
855                        return Err(XlogError::Execution(format!(
856                            "Column index {} out of bounds",
857                            col_idx
858                        )));
859                    }
860                }
861                ProjectExpr::Computed(_expr, result_type) => {
862                    // Computed columns get a generated name
863                    let col_name = format!("computed_{}", projected_columns.len());
864                    projected_columns.push((col_name, *result_type));
865                    projected_sort_labels.push(format!("computed_{}", projected_sort_labels.len()));
866                }
867            }
868        }
869        Schema::new(projected_columns)
870            .with_sort_labels(projected_sort_labels)
871            .map_err(XlogError::Execution)
872    }
873}