1use 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 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 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 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 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 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 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 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 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 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 pub(crate) fn execute_project(
780 &self,
781 input: &CudaBuffer,
782 columns: &[ProjectExpr],
783 ) -> Result<CudaBuffer> {
784 if input.is_empty() {
785 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 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 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 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 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 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 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}