micromegas_datafusion_extensions/jsonb/
each.rs1use 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#[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#[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
86fn 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
144fn 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#[derive(Debug)]
192pub struct JsonbEachTableProvider {
193 source: JsonbSource,
194}
195
196impl JsonbEachTableProvider {
197 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 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}