Skip to main content

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