1use std::hash::{Hash, Hasher};
18use std::sync::Arc;
19
20use arrow::datatypes::DataType;
21use arrow_array::{Array, ArrayRef, BooleanArray, ListArray, MapArray, StructArray};
22use arrow_schema::{Field, Fields, Schema};
23use datafusion_common::Result;
24use datafusion_expr::{
25 ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, TypeSignature, Volatility,
26};
27
28use super::json_udf::expand_args;
29
30pub fn register_lambda_functions(ctx: &datafusion::prelude::SessionContext) {
32 use datafusion_expr::ScalarUDF;
33
34 ctx.register_udf(ScalarUDF::new_from_impl(ArrayTransform::new()));
35 ctx.register_udf(ScalarUDF::new_from_impl(ArrayFilter::new()));
36 ctx.register_udf(ScalarUDF::new_from_impl(ArrayReduce::new()));
37 ctx.register_udf(ScalarUDF::new_from_impl(MapFilter::new()));
38 ctx.register_udf(ScalarUDF::new_from_impl(MapTransformValues::new()));
39}
40
41thread_local! {
42 static LAMBDA_CTX: std::cell::RefCell<Option<datafusion::prelude::SessionContext>> =
48 const { std::cell::RefCell::new(None) };
49}
50
51fn eval_expr_on_batch(sql_expr: &str, batch: &arrow_array::RecordBatch) -> Result<ArrayRef> {
55 let ctx = LAMBDA_CTX.with(|cell| {
58 let mut opt = cell.borrow_mut();
59 opt.get_or_insert_with(datafusion::prelude::SessionContext::new)
60 .clone()
61 });
62
63 let provider =
64 datafusion::datasource::MemTable::try_new(batch.schema(), vec![vec![batch.clone()]])?;
65 let rt = tokio::runtime::Handle::try_current().map_err(|e| {
66 datafusion_common::DataFusionError::Internal(format!(
67 "lambda eval requires tokio runtime: {e}"
68 ))
69 })?;
70 tokio::task::block_in_place(|| {
72 rt.block_on(async {
73 ctx.register_table("__lambda_data", Arc::new(provider))?;
74 let df = ctx
75 .sql(&format!("SELECT {sql_expr} FROM __lambda_data"))
76 .await?;
77 let batches = df.collect().await?;
78 let _ = ctx.deregister_table("__lambda_data");
80 if batches.is_empty() {
81 Err(datafusion_common::DataFusionError::Internal(
82 "lambda expression returned no data".into(),
83 ))
84 } else {
85 let result = arrow::compute::concat_batches(&batches[0].schema(), &batches)?;
87 Ok(result.column(0).clone())
88 }
89 })
90 })
91}
92
93fn scalar_string_value(cv: &ColumnarValue) -> Result<String> {
94 match cv {
95 ColumnarValue::Scalar(s) => {
96 let arr = s.to_array_of_size(1)?;
97 let str_arr = arr
98 .as_any()
99 .downcast_ref::<arrow_array::StringArray>()
100 .ok_or_else(|| {
101 datafusion_common::DataFusionError::Internal("expected Utf8 argument".into())
102 })?;
103 Ok(str_arr.value(0).to_string())
104 }
105 ColumnarValue::Array(arr) => {
106 let str_arr = arr
107 .as_any()
108 .downcast_ref::<arrow_array::StringArray>()
109 .ok_or_else(|| {
110 datafusion_common::DataFusionError::Internal("expected Utf8 argument".into())
111 })?;
112 Ok(str_arr.value(0).to_string())
113 }
114 }
115}
116
117#[derive(Debug)]
123pub struct ArrayTransform {
124 signature: Signature,
125}
126
127impl ArrayTransform {
128 #[must_use]
130 pub fn new() -> Self {
131 Self {
132 signature: Signature::new(TypeSignature::Any(2), Volatility::Immutable),
133 }
134 }
135}
136
137impl Default for ArrayTransform {
138 fn default() -> Self {
139 Self::new()
140 }
141}
142impl PartialEq for ArrayTransform {
143 fn eq(&self, _: &Self) -> bool {
144 true
145 }
146}
147impl Eq for ArrayTransform {}
148impl Hash for ArrayTransform {
149 fn hash<H: Hasher>(&self, s: &mut H) {
150 "array_transform".hash(s);
151 }
152}
153
154impl ScalarUDFImpl for ArrayTransform {
155 fn as_any(&self) -> &dyn std::any::Any {
156 self
157 }
158
159 fn name(&self) -> &'static str {
160 "array_transform"
161 }
162 fn signature(&self) -> &Signature {
163 &self.signature
164 }
165
166 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
167 match &arg_types[0] {
168 DataType::List(f) => Ok(DataType::List(Arc::clone(f))),
169 _ => Ok(DataType::List(Arc::new(Field::new(
170 "item",
171 DataType::Utf8,
172 true,
173 )))),
174 }
175 }
176
177 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
178 let expanded = expand_args(&args.args)?;
179 let list_arr = expanded[0]
180 .as_any()
181 .downcast_ref::<ListArray>()
182 .ok_or_else(|| {
183 datafusion_common::DataFusionError::Internal(
184 "array_transform: first arg must be List".into(),
185 )
186 })?;
187
188 let lambda_str = scalar_string_value(&args.args[1])?;
189 let flat_values = list_arr.values();
190
191 let schema = Arc::new(Schema::new(vec![Field::new(
192 "x",
193 flat_values.data_type().clone(),
194 true,
195 )]));
196 let batch = arrow_array::RecordBatch::try_new(schema, vec![Arc::clone(flat_values)])?;
197
198 let result_arr = eval_expr_on_batch(&lambda_str, &batch)?;
199
200 let new_field = Arc::new(Field::new("item", result_arr.data_type().clone(), true));
201 let new_list = ListArray::try_new(
202 new_field,
203 list_arr.offsets().clone(),
204 result_arr,
205 list_arr.nulls().cloned(),
206 )?;
207 Ok(ColumnarValue::Array(Arc::new(new_list)))
208 }
209}
210
211#[derive(Debug)]
217pub struct ArrayFilter {
218 signature: Signature,
219}
220
221impl ArrayFilter {
222 #[must_use]
224 pub fn new() -> Self {
225 Self {
226 signature: Signature::new(TypeSignature::Any(2), Volatility::Immutable),
227 }
228 }
229}
230
231impl Default for ArrayFilter {
232 fn default() -> Self {
233 Self::new()
234 }
235}
236impl PartialEq for ArrayFilter {
237 fn eq(&self, _: &Self) -> bool {
238 true
239 }
240}
241impl Eq for ArrayFilter {}
242impl Hash for ArrayFilter {
243 fn hash<H: Hasher>(&self, s: &mut H) {
244 "array_filter".hash(s);
245 }
246}
247
248impl ScalarUDFImpl for ArrayFilter {
249 fn as_any(&self) -> &dyn std::any::Any {
250 self
251 }
252
253 fn name(&self) -> &'static str {
254 "array_filter"
255 }
256 fn signature(&self) -> &Signature {
257 &self.signature
258 }
259
260 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
261 match &arg_types[0] {
262 DataType::List(f) => Ok(DataType::List(Arc::clone(f))),
263 _ => Ok(DataType::List(Arc::new(Field::new(
264 "item",
265 DataType::Utf8,
266 true,
267 )))),
268 }
269 }
270
271 #[allow(
272 clippy::cast_sign_loss,
273 clippy::cast_possible_wrap,
274 clippy::cast_possible_truncation
275 )]
276 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
277 let expanded = expand_args(&args.args)?;
278 let list_arr = expanded[0]
279 .as_any()
280 .downcast_ref::<ListArray>()
281 .ok_or_else(|| {
282 datafusion_common::DataFusionError::Internal(
283 "array_filter: first arg must be List".into(),
284 )
285 })?;
286
287 let lambda_str = scalar_string_value(&args.args[1])?;
288 let flat_values = list_arr.values();
289 let elem_type = flat_values.data_type().clone();
290
291 let schema = Arc::new(Schema::new(vec![Field::new("x", elem_type.clone(), true)]));
292 let batch = arrow_array::RecordBatch::try_new(schema, vec![Arc::clone(flat_values)])?;
293
294 let mask_arr = eval_expr_on_batch(&lambda_str, &batch)?;
295 let mask = mask_arr
296 .as_any()
297 .downcast_ref::<BooleanArray>()
298 .ok_or_else(|| {
299 datafusion_common::DataFusionError::Internal(
300 "array_filter: lambda must return Boolean".into(),
301 )
302 })?;
303
304 let mut offsets = vec![0i32];
305 let mut filtered_indices: Vec<usize> = Vec::new();
306
307 for row in 0..list_arr.len() {
308 let start = list_arr.value_offsets()[row] as usize;
309 let end = list_arr.value_offsets()[row + 1] as usize;
310
311 for i in start..end {
312 if !mask.is_null(i) && mask.value(i) {
313 filtered_indices.push(i);
314 }
315 }
316 offsets.push(filtered_indices.len() as i32);
317 }
318
319 let indices = arrow_array::UInt32Array::from(
320 filtered_indices
321 .iter()
322 .map(|&i| i as u32)
323 .collect::<Vec<_>>(),
324 );
325 let filtered_values = arrow::compute::take(flat_values.as_ref(), &indices, None)?;
326
327 let new_field = Arc::new(Field::new("item", elem_type, true));
328 let new_offsets =
329 arrow::buffer::OffsetBuffer::new(arrow::buffer::ScalarBuffer::from(offsets));
330 let new_list = ListArray::try_new(
331 new_field,
332 new_offsets,
333 filtered_values,
334 list_arr.nulls().cloned(),
335 )?;
336 Ok(ColumnarValue::Array(Arc::new(new_list)))
337 }
338}
339
340#[derive(Debug)]
346pub struct ArrayReduce {
347 signature: Signature,
348}
349
350impl ArrayReduce {
351 #[must_use]
353 pub fn new() -> Self {
354 Self {
355 signature: Signature::new(TypeSignature::Any(3), Volatility::Immutable),
356 }
357 }
358}
359
360impl Default for ArrayReduce {
361 fn default() -> Self {
362 Self::new()
363 }
364}
365impl PartialEq for ArrayReduce {
366 fn eq(&self, _: &Self) -> bool {
367 true
368 }
369}
370impl Eq for ArrayReduce {}
371impl Hash for ArrayReduce {
372 fn hash<H: Hasher>(&self, s: &mut H) {
373 "array_reduce".hash(s);
374 }
375}
376
377impl ScalarUDFImpl for ArrayReduce {
378 fn as_any(&self) -> &dyn std::any::Any {
379 self
380 }
381
382 fn name(&self) -> &'static str {
383 "array_reduce"
384 }
385 fn signature(&self) -> &Signature {
386 &self.signature
387 }
388
389 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
390 Ok(arg_types.get(1).cloned().unwrap_or(DataType::Int64))
391 }
392
393 #[allow(clippy::cast_sign_loss)]
394 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
395 let expanded = expand_args(&args.args)?;
396 let list_arr = expanded[0]
397 .as_any()
398 .downcast_ref::<ListArray>()
399 .ok_or_else(|| {
400 datafusion_common::DataFusionError::Internal(
401 "array_reduce: first arg must be List".into(),
402 )
403 })?;
404
405 let init_arr = &expanded[1];
406 let lambda_str = scalar_string_value(&args.args[2])?;
407
408 let elem_type = list_arr.values().data_type().clone();
409 let acc_type = init_arr.data_type().clone();
410
411 let schema = Arc::new(Schema::new(vec![
412 Field::new("acc", acc_type, true),
413 Field::new("x", elem_type, true),
414 ]));
415
416 let mut result_builder: Vec<ArrayRef> = Vec::new();
417
418 for row in 0..list_arr.len() {
419 let start = list_arr.value_offsets()[row] as usize;
420 let end = list_arr.value_offsets()[row + 1] as usize;
421
422 let mut acc: ArrayRef = init_arr.slice(row, 1);
423
424 for i in start..end {
425 let x = list_arr.values().slice(i, 1);
426 let batch = arrow_array::RecordBatch::try_new(
427 Arc::clone(&schema),
428 vec![Arc::clone(&acc), x],
429 )?;
430 let result_col = eval_expr_on_batch(&lambda_str, &batch)?;
431 acc = result_col;
432 }
433
434 result_builder.push(acc);
435 }
436
437 if result_builder.is_empty() {
438 return Ok(ColumnarValue::Array(Arc::clone(init_arr)));
439 }
440
441 let refs: Vec<&dyn Array> = result_builder
442 .iter()
443 .map(std::convert::AsRef::as_ref)
444 .collect();
445 let result = arrow::compute::concat(&refs)?;
446 Ok(ColumnarValue::Array(result))
447 }
448}
449
450#[derive(Debug)]
456pub struct MapFilter {
457 signature: Signature,
458}
459
460impl MapFilter {
461 #[must_use]
463 pub fn new() -> Self {
464 Self {
465 signature: Signature::new(TypeSignature::Any(2), Volatility::Immutable),
466 }
467 }
468}
469
470impl Default for MapFilter {
471 fn default() -> Self {
472 Self::new()
473 }
474}
475impl PartialEq for MapFilter {
476 fn eq(&self, _: &Self) -> bool {
477 true
478 }
479}
480impl Eq for MapFilter {}
481impl Hash for MapFilter {
482 fn hash<H: Hasher>(&self, s: &mut H) {
483 "map_filter".hash(s);
484 }
485}
486
487impl ScalarUDFImpl for MapFilter {
488 fn as_any(&self) -> &dyn std::any::Any {
489 self
490 }
491
492 fn name(&self) -> &'static str {
493 "map_filter"
494 }
495 fn signature(&self) -> &Signature {
496 &self.signature
497 }
498
499 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
500 Ok(arg_types[0].clone())
501 }
502
503 #[allow(
504 clippy::cast_sign_loss,
505 clippy::cast_possible_wrap,
506 clippy::cast_possible_truncation
507 )]
508 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
509 let expanded = expand_args(&args.args)?;
510 let map_arr = expanded[0]
511 .as_any()
512 .downcast_ref::<MapArray>()
513 .ok_or_else(|| {
514 datafusion_common::DataFusionError::Internal(
515 "map_filter: first arg must be Map".into(),
516 )
517 })?;
518
519 let lambda_str = scalar_string_value(&args.args[1])?;
520
521 let entries = map_arr.entries();
522 let key_col = entries.column(0);
523 let val_col = entries.column(1);
524
525 let key_type = key_col.data_type().clone();
526 let val_type = val_col.data_type().clone();
527
528 let schema = Arc::new(Schema::new(vec![
529 Field::new("k", key_type.clone(), true),
530 Field::new("v", val_type.clone(), true),
531 ]));
532 let batch = arrow_array::RecordBatch::try_new(
533 schema,
534 vec![Arc::clone(key_col), Arc::clone(val_col)],
535 )?;
536
537 let mask_arr = eval_expr_on_batch(&lambda_str, &batch)?;
538 let mask_bool = mask_arr
539 .as_any()
540 .downcast_ref::<BooleanArray>()
541 .ok_or_else(|| {
542 datafusion_common::DataFusionError::Internal(
543 "map_filter: lambda must return Boolean".into(),
544 )
545 })?;
546
547 let mut offsets = vec![0i32];
548 let mut keep_indices: Vec<usize> = Vec::new();
549
550 for row in 0..map_arr.len() {
551 let start = map_arr.value_offsets()[row] as usize;
552 let end = map_arr.value_offsets()[row + 1] as usize;
553
554 for i in start..end {
555 if !mask_bool.is_null(i) && mask_bool.value(i) {
556 keep_indices.push(i);
557 }
558 }
559 offsets.push(keep_indices.len() as i32);
560 }
561
562 let indices = arrow_array::UInt32Array::from(
563 keep_indices.iter().map(|&i| i as u32).collect::<Vec<_>>(),
564 );
565 let new_keys = arrow::compute::take(key_col.as_ref(), &indices, None)?;
566 let new_vals = arrow::compute::take(val_col.as_ref(), &indices, None)?;
567
568 let struct_fields = Fields::from(vec![
569 Field::new("key", key_type, false),
570 Field::new("value", val_type, true),
571 ]);
572 let new_entries = StructArray::try_new(struct_fields, vec![new_keys, new_vals], None)?;
573
574 let entries_field = Field::new("entries", new_entries.data_type().clone(), false);
575 let new_offsets =
576 arrow::buffer::OffsetBuffer::new(arrow::buffer::ScalarBuffer::from(offsets));
577 let new_map = MapArray::try_new(
578 Arc::new(entries_field),
579 new_offsets,
580 new_entries,
581 map_arr.nulls().cloned(),
582 false,
583 )?;
584 Ok(ColumnarValue::Array(Arc::new(new_map)))
585 }
586}
587
588#[derive(Debug)]
594pub struct MapTransformValues {
595 signature: Signature,
596}
597
598impl MapTransformValues {
599 #[must_use]
601 pub fn new() -> Self {
602 Self {
603 signature: Signature::new(TypeSignature::Any(2), Volatility::Immutable),
604 }
605 }
606}
607
608impl Default for MapTransformValues {
609 fn default() -> Self {
610 Self::new()
611 }
612}
613impl PartialEq for MapTransformValues {
614 fn eq(&self, _: &Self) -> bool {
615 true
616 }
617}
618impl Eq for MapTransformValues {}
619impl Hash for MapTransformValues {
620 fn hash<H: Hasher>(&self, s: &mut H) {
621 "map_transform_values".hash(s);
622 }
623}
624
625impl ScalarUDFImpl for MapTransformValues {
626 fn as_any(&self) -> &dyn std::any::Any {
627 self
628 }
629
630 fn name(&self) -> &'static str {
631 "map_transform_values"
632 }
633 fn signature(&self) -> &Signature {
634 &self.signature
635 }
636
637 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
638 Ok(arg_types[0].clone())
639 }
640
641 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
642 let expanded = expand_args(&args.args)?;
643 let map_arr = expanded[0]
644 .as_any()
645 .downcast_ref::<MapArray>()
646 .ok_or_else(|| {
647 datafusion_common::DataFusionError::Internal(
648 "map_transform_values: first arg must be Map".into(),
649 )
650 })?;
651
652 let lambda_str = scalar_string_value(&args.args[1])?;
653
654 let entries = map_arr.entries();
655 let key_col = entries.column(0);
656 let val_col = entries.column(1);
657
658 let key_type = key_col.data_type().clone();
659
660 let schema = Arc::new(Schema::new(vec![
661 Field::new("k", key_type.clone(), true),
662 Field::new("v", val_col.data_type().clone(), true),
663 ]));
664 let batch = arrow_array::RecordBatch::try_new(
665 schema,
666 vec![Arc::clone(key_col), Arc::clone(val_col)],
667 )?;
668
669 let new_vals = eval_expr_on_batch(&lambda_str, &batch)?;
670
671 let struct_fields = Fields::from(vec![
672 Field::new("key", key_type, false),
673 Field::new("value", new_vals.data_type().clone(), true),
674 ]);
675 let new_entries =
676 StructArray::try_new(struct_fields, vec![Arc::clone(key_col), new_vals], None)?;
677
678 let entries_field = Field::new("entries", new_entries.data_type().clone(), false);
679 let new_map = MapArray::try_new(
680 Arc::new(entries_field),
681 map_arr.offsets().clone(),
682 new_entries,
683 map_arr.nulls().cloned(),
684 false,
685 )?;
686 Ok(ColumnarValue::Array(Arc::new(new_map)))
687 }
688}
689
690#[cfg(test)]
691mod tests {
692 use super::*;
693 use crate::datafusion::create_session_context;
694 use arrow_array::*;
695 use datafusion_common::config::ConfigOptions;
696
697 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
700 async fn test_array_transform_add_one() {
701 let values = Int64Array::from(vec![1, 2, 3, 4, 5, 6]);
702 let offsets =
703 arrow::buffer::OffsetBuffer::new(arrow::buffer::ScalarBuffer::from(vec![0i32, 3, 6]));
704 let list = ListArray::try_new(
705 Arc::new(Field::new("item", DataType::Int64, true)),
706 offsets,
707 Arc::new(values),
708 None,
709 )
710 .unwrap();
711
712 let udf = ArrayTransform::new();
713 let result = udf
714 .invoke_with_args(ScalarFunctionArgs {
715 args: vec![
716 ColumnarValue::Array(Arc::new(list)),
717 ColumnarValue::Scalar(datafusion_common::ScalarValue::Utf8(Some(
718 "x + 1".into(),
719 ))),
720 ],
721 number_rows: 0,
722 arg_fields: vec![],
723 return_field: Arc::new(Field::new(
724 "output",
725 DataType::List(Arc::new(Field::new("item", DataType::Int64, true))),
726 true,
727 )),
728 config_options: Arc::new(ConfigOptions::default()),
729 })
730 .unwrap();
731
732 if let ColumnarValue::Array(arr) = result {
733 let la = arr.as_any().downcast_ref::<ListArray>().unwrap();
734 assert_eq!(la.len(), 2);
735 let row0 = la.value(0);
736 let r0 = row0.as_any().downcast_ref::<Int64Array>().unwrap();
737 assert_eq!(r0.value(0), 2);
738 assert_eq!(r0.value(1), 3);
739 assert_eq!(r0.value(2), 4);
740 } else {
741 panic!("expected Array");
742 }
743 }
744
745 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
748 async fn test_array_filter_positive() {
749 let values = Int64Array::from(vec![-1, 2, -3, 4]);
750 let offsets =
751 arrow::buffer::OffsetBuffer::new(arrow::buffer::ScalarBuffer::from(vec![0i32, 4]));
752 let list = ListArray::try_new(
753 Arc::new(Field::new("item", DataType::Int64, true)),
754 offsets,
755 Arc::new(values),
756 None,
757 )
758 .unwrap();
759
760 let udf = ArrayFilter::new();
761 let result = udf
762 .invoke_with_args(ScalarFunctionArgs {
763 args: vec![
764 ColumnarValue::Array(Arc::new(list)),
765 ColumnarValue::Scalar(datafusion_common::ScalarValue::Utf8(Some(
766 "x > 0".into(),
767 ))),
768 ],
769 number_rows: 0,
770 arg_fields: vec![],
771 return_field: Arc::new(Field::new(
772 "output",
773 DataType::List(Arc::new(Field::new("item", DataType::Int64, true))),
774 true,
775 )),
776 config_options: Arc::new(ConfigOptions::default()),
777 })
778 .unwrap();
779
780 if let ColumnarValue::Array(arr) = result {
781 let la = arr.as_any().downcast_ref::<ListArray>().unwrap();
782 let row0 = la.value(0);
783 let r0 = row0.as_any().downcast_ref::<Int64Array>().unwrap();
784 assert_eq!(r0.len(), 2);
785 assert_eq!(r0.value(0), 2);
786 assert_eq!(r0.value(1), 4);
787 }
788 }
789
790 #[test]
793 fn test_register_lambda_functions() {
794 use datafusion::execution::FunctionRegistry;
795
796 let ctx = create_session_context();
797 register_lambda_functions(&ctx);
798 assert!(ctx.udf("array_transform").is_ok());
799 assert!(ctx.udf("array_filter").is_ok());
800 assert!(ctx.udf("array_reduce").is_ok());
801 assert!(ctx.udf("map_filter").is_ok());
802 assert!(ctx.udf("map_transform_values").is_ok());
803 }
804}