Skip to main content

micromegas_datafusion_extensions/jsonb/
array_elements.rs

1use async_trait::async_trait;
2use datafusion::arrow::array::{Array, ArrayRef, BinaryArray, DictionaryArray, GenericBinaryArray};
3use datafusion::arrow::datatypes::{DataType, Field, Int32Type, Schema, SchemaRef};
4use datafusion::arrow::record_batch::RecordBatch;
5use datafusion::catalog::Session;
6use datafusion::catalog::TableFunctionArgs;
7use datafusion::catalog::TableFunctionImpl;
8use datafusion::catalog::TableProvider;
9use datafusion::datasource::TableType;
10use datafusion::datasource::memory::{DataSourceExec, MemorySourceConfig};
11use datafusion::error::DataFusionError;
12use datafusion::logical_expr::{LogicalPlan, LogicalPlanBuilder};
13use datafusion::physical_plan::ExecutionPlan;
14use datafusion::prelude::Expr;
15use datafusion::scalar::ScalarValue;
16use jsonb::RawJsonb;
17use std::sync::Arc;
18
19/// A DataFusion `TableFunctionImpl` that expands a JSONB array into rows with a single `value` column.
20///
21/// Usage:
22/// ```sql
23/// SELECT jsonb_as_string(elem.value)
24/// FROM jsonb_array_elements(jsonb_parse('[1, 2, 3]')) as elem
25/// ```
26#[derive(Debug)]
27pub struct JsonbArrayElementsTableFunction {}
28
29impl JsonbArrayElementsTableFunction {
30    pub fn new() -> Self {
31        Self {}
32    }
33}
34
35impl Default for JsonbArrayElementsTableFunction {
36    fn default() -> Self {
37        Self::new()
38    }
39}
40
41/// The source of JSONB data — either a literal value or a subquery/expression to evaluate.
42#[derive(Debug, Clone)]
43enum JsonbSource {
44    Literal(ScalarValue),
45    Subquery(Arc<LogicalPlan>),
46}
47
48impl TableFunctionImpl for JsonbArrayElementsTableFunction {
49    fn call_with_args(
50        &self,
51        args: TableFunctionArgs,
52    ) -> datafusion::error::Result<Arc<dyn TableProvider>> {
53        let args = args.exprs();
54        if args.len() != 1 {
55            return Err(DataFusionError::Plan(
56                "jsonb_array_elements requires exactly one argument (a JSONB array)".into(),
57            ));
58        }
59
60        let source = match &args[0] {
61            Expr::Literal(scalar, _metadata) => JsonbSource::Literal(scalar.clone()),
62            Expr::ScalarSubquery(subquery) => JsonbSource::Subquery(subquery.subquery.clone()),
63            other => {
64                let plan = LogicalPlanBuilder::empty(true)
65                    .project(vec![other.clone()])?
66                    .build()?;
67                JsonbSource::Subquery(Arc::new(plan))
68            }
69        };
70
71        Ok(Arc::new(JsonbArrayElementsTableProvider { source }))
72    }
73}
74
75fn output_schema() -> SchemaRef {
76    Arc::new(Schema::new(vec![Field::new(
77        "value",
78        DataType::Binary,
79        false,
80    )]))
81}
82
83/// Extract element values from a JSONB array.
84fn extract_elements_from_jsonb(jsonb_bytes: &[u8]) -> Result<Vec<Vec<u8>>, DataFusionError> {
85    let jsonb = RawJsonb::new(jsonb_bytes);
86    match jsonb.array_values() {
87        Ok(Some(values)) => Ok(values.into_iter().map(|v| v.as_ref().to_vec()).collect()),
88        Ok(None) => Err(DataFusionError::Execution(
89            "jsonb_array_elements: input is not a JSONB array".into(),
90        )),
91        Err(e) => Err(DataFusionError::External(e.into())),
92    }
93}
94
95fn empty_exec(projection: Option<&Vec<usize>>) -> Result<Arc<dyn ExecutionPlan>, DataFusionError> {
96    let batch = RecordBatch::new_empty(output_schema());
97    let source = MemorySourceConfig::try_new(
98        &[vec![batch]],
99        output_schema(),
100        projection.map(|v| v.to_owned()),
101    )?;
102    Ok(DataSourceExec::from_data_source(source))
103}
104
105fn elements_to_batch(elements: &[Vec<u8>]) -> Result<RecordBatch, DataFusionError> {
106    if elements.is_empty() {
107        return Ok(RecordBatch::new_empty(output_schema()));
108    }
109
110    let values: Vec<&[u8]> = elements.iter().map(|v| v.as_slice()).collect();
111    let value_array: ArrayRef = Arc::new(BinaryArray::from(values));
112
113    RecordBatch::try_new(output_schema(), vec![value_array])
114        .map_err(|e| DataFusionError::External(e.into()))
115}
116
117fn scalar_to_elements(scalar: &ScalarValue) -> Result<Vec<Vec<u8>>, DataFusionError> {
118    match scalar {
119        ScalarValue::Binary(Some(bytes)) => extract_elements_from_jsonb(bytes),
120        ScalarValue::Binary(None) => Ok(vec![]),
121        ScalarValue::Dictionary(_, inner) => scalar_to_elements(inner.as_ref()),
122        _ => Err(DataFusionError::Plan(format!(
123            "jsonb_array_elements argument must be Binary (JSONB), got: {:?}",
124            scalar.data_type()
125        ))),
126    }
127}
128
129/// Extract JSONB bytes from all rows of a column, handling both plain Binary
130/// and Dictionary<Int32, Binary> encodings.
131fn extract_all_jsonb_bytes_from_column(column: &ArrayRef) -> Result<Vec<Vec<u8>>, DataFusionError> {
132    match column.data_type() {
133        DataType::Binary => {
134            let binary_array = column
135                .as_any()
136                .downcast_ref::<GenericBinaryArray<i32>>()
137                .ok_or_else(|| {
138                    DataFusionError::Execution("failed to cast column to BinaryArray".into())
139                })?;
140            Ok((0..binary_array.len())
141                .filter(|&i| !binary_array.is_null(i))
142                .map(|i| binary_array.value(i).to_vec())
143                .collect())
144        }
145        DataType::Dictionary(_, value_type) if matches!(value_type.as_ref(), DataType::Binary) => {
146            let dict_array = column
147                .as_any()
148                .downcast_ref::<DictionaryArray<Int32Type>>()
149                .ok_or_else(|| {
150                    DataFusionError::Execution(
151                        "failed to cast column to DictionaryArray<Int32, Binary>".into(),
152                    )
153                })?;
154            let binary_values = dict_array
155                .values()
156                .as_any()
157                .downcast_ref::<GenericBinaryArray<i32>>()
158                .ok_or_else(|| {
159                    DataFusionError::Execution("dictionary values are not a binary array".into())
160                })?;
161            Ok((0..dict_array.len())
162                .filter(|&i| !dict_array.is_null(i))
163                .map(|i| {
164                    let key_index = dict_array.keys().value(i) as usize;
165                    binary_values.value(key_index).to_vec()
166                })
167                .collect())
168        }
169        other => Err(DataFusionError::Execution(format!(
170            "jsonb_array_elements subquery must return a Binary or Dictionary<Int32, Binary> column, got: {other:?}"
171        ))),
172    }
173}
174
175/// Table provider for expanding JSONB arrays into value rows.
176#[derive(Debug)]
177pub struct JsonbArrayElementsTableProvider {
178    source: JsonbSource,
179}
180
181impl JsonbArrayElementsTableProvider {
182    /// Creates a new provider from a JSONB scalar value (for testing).
183    pub fn from_scalar(scalar: ScalarValue) -> Result<Self, DataFusionError> {
184        if !matches!(&scalar, ScalarValue::Binary(Some(_))) {
185            return Err(DataFusionError::Plan(format!(
186                "jsonb_array_elements argument must be Binary (JSONB), got: {:?}",
187                scalar.data_type()
188            )));
189        }
190        Ok(Self {
191            source: JsonbSource::Literal(scalar),
192        })
193    }
194}
195
196#[async_trait]
197impl TableProvider for JsonbArrayElementsTableProvider {
198    fn schema(&self) -> SchemaRef {
199        output_schema()
200    }
201
202    fn table_type(&self) -> TableType {
203        TableType::Temporary
204    }
205
206    async fn scan(
207        &self,
208        state: &dyn Session,
209        projection: Option<&Vec<usize>>,
210        _filters: &[Expr],
211        limit: Option<usize>,
212    ) -> datafusion::error::Result<Arc<dyn ExecutionPlan>> {
213        let elements = match &self.source {
214            JsonbSource::Literal(scalar) => scalar_to_elements(scalar)?,
215            JsonbSource::Subquery(plan) => {
216                let physical_plan = state.create_physical_plan(plan).await?;
217                let task_ctx = state.task_ctx();
218                let batches = datafusion::physical_plan::collect(physical_plan, task_ctx).await?;
219
220                let mut all_elements = Vec::new();
221                if batches.is_empty() || batches.iter().all(|b| b.num_rows() == 0) {
222                    return empty_exec(projection);
223                }
224                for batch in &batches {
225                    if batch.num_columns() != 1 {
226                        return Err(DataFusionError::Execution(format!(
227                            "jsonb_array_elements subquery must return exactly one column, got {}",
228                            batch.num_columns()
229                        )));
230                    }
231                    for jsonb_bytes in extract_all_jsonb_bytes_from_column(batch.column(0))? {
232                        all_elements.extend(extract_elements_from_jsonb(&jsonb_bytes)?);
233                    }
234                }
235                all_elements
236            }
237        };
238
239        let mut record_batch = elements_to_batch(&elements)?;
240
241        // Apply limit if specified
242        if let Some(n) = limit
243            && n < record_batch.num_rows()
244        {
245            record_batch = record_batch.slice(0, n);
246        }
247
248        let source = MemorySourceConfig::try_new(
249            &[vec![record_batch]],
250            self.schema(),
251            projection.map(|v| v.to_owned()),
252        )?;
253        Ok(DataSourceExec::from_data_source(source))
254    }
255}