Skip to main content

laminar_sql/datafusion/
complex_type_lambda.rs

1//! Lambda higher-order functions for arrays and maps (F-SCHEMA-015 Tier 3).
2//!
3//! Provides vectorized lambda evaluation over Arrow arrays:
4//!
5//! | Function | Lambda | Strategy |
6//! |----------|--------|----------|
7//! | `array_transform(arr, lambda)` | `x -> expr` | Flatten, eval, re-group |
8//! | `array_filter(arr, lambda)` | `x -> bool` | Flatten, eval, filter+rebuild |
9//! | `array_reduce(arr, init, lambda)` | `(acc, x) -> expr` | Sequential fold |
10//! | `map_filter(map, lambda)` | `(k, v) -> bool` | Eval on k+v, filter entries |
11//! | `map_transform_values(map, lambda)` | `(k, v) -> expr` | Eval on k+v, replace vals |
12//!
13//! Lambda expressions are specified as string literal SQL expressions.
14//! They are evaluated using DataFusion's SQL engine against a temporary
15//! table containing the element values. Native lambda syntax is deferred.
16
17use 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
30/// Registers all lambda HOFs with the given session context.
31pub 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    /// Cached `SessionContext` for lambda evaluation.
43    ///
44    /// Creating a `SessionContext` is expensive because it registers all
45    /// built-in functions and sets up catalogs. We cache it per-thread
46    /// and reuse it across lambda invocations.
47    static LAMBDA_CTX: std::cell::RefCell<Option<datafusion::prelude::SessionContext>> =
48        const { std::cell::RefCell::new(None) };
49}
50
51/// Evaluate a SQL expression against a `RecordBatch`, returning the result column.
52///
53/// The expression can reference columns by name from the batch schema.
54fn eval_expr_on_batch(sql_expr: &str, batch: &arrow_array::RecordBatch) -> Result<ArrayRef> {
55    // Reuse a thread-local SessionContext to avoid the cost of
56    // registering built-in functions on every invocation.
57    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    // Use block_in_place to allow blocking inside an already-running tokio runtime.
71    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            // Deregister the ephemeral table so the batch data is freed.
79            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                // Concatenate all result batches and return the first column.
86                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// ══════════════════════════════════════════════════════════════════
118// array_transform(arr, lambda_str) -> List
119// ══════════════════════════════════════════════════════════════════
120
121/// `array_transform(arr, 'x + 1')` — apply a lambda to each element.
122#[derive(Debug)]
123pub struct ArrayTransform {
124    signature: Signature,
125}
126
127impl ArrayTransform {
128    /// Creates a new `array_transform` UDF.
129    #[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// ══════════════════════════════════════════════════════════════════
212// array_filter(arr, lambda_str) -> List
213// ══════════════════════════════════════════════════════════════════
214
215/// `array_filter(arr, 'x > 0')` — filter elements by a boolean lambda.
216#[derive(Debug)]
217pub struct ArrayFilter {
218    signature: Signature,
219}
220
221impl ArrayFilter {
222    /// Creates a new `array_filter` UDF.
223    #[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// ══════════════════════════════════════════════════════════════════
341// array_reduce(arr, init, lambda_str) -> scalar
342// ══════════════════════════════════════════════════════════════════
343
344/// `array_reduce(arr, init, '(acc + x)')` — fold/reduce array elements.
345#[derive(Debug)]
346pub struct ArrayReduce {
347    signature: Signature,
348}
349
350impl ArrayReduce {
351    /// Creates a new `array_reduce` UDF.
352    #[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// ══════════════════════════════════════════════════════════════════
451// map_filter(map, lambda_str) -> Map
452// ══════════════════════════════════════════════════════════════════
453
454/// `map_filter(map, '(k <> ''temp'')')` — filter map entries by key+value.
455#[derive(Debug)]
456pub struct MapFilter {
457    signature: Signature,
458}
459
460impl MapFilter {
461    /// Creates a new `map_filter` UDF.
462    #[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// ══════════════════════════════════════════════════════════════════
589// map_transform_values(map, lambda_str) -> Map
590// ══════════════════════════════════════════════════════════════════
591
592/// `map_transform_values(map, 'v * 2')` — transform map values.
593#[derive(Debug)]
594pub struct MapTransformValues {
595    signature: Signature,
596}
597
598impl MapTransformValues {
599    /// Creates a new `map_transform_values` UDF.
600    #[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    // ── array_transform ─────────────────────────────────────────
698
699    #[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    // ── array_filter ────────────────────────────────────────────
746
747    #[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    // ── Registration ────────────────────────────────────────────
791
792    #[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}