Skip to content
Prev Previous commit
Next Next commit
feat(udf): preserve scalar-vs-array across JNI in invoke_with_args
JavaScalarUdf::invoke_with_args now partitions ScalarFunctionArgs::args
by ColumnarValue variant: array args travel through a length-N struct,
scalar args through a separate length-1 struct (one element each), and
a positional byte slice tells the Java side how to interleave them
back. The JNI signature of invokeScalarUdf becomes
(Lorg/apache/datafusion/ScalarFunction;JJJJ[BJJI)B; the returned byte
indicates Array (0) or Scalar (1) so the native side reconstructs the
right ColumnarValue variant via ScalarValue::try_from_array.

Drops the prior scalar-materialisation step, which was the workaround
that PR #57 attempted to patch by passing rowCount; nullary UDFs that
broadcast a value now return ColumnarValue.Scalar instead.

Closes #62.
  • Loading branch information
andygrove committed May 18, 2026
commit 632a7e0c7ef8852979f8df356f7f5eb65bc133c2
2 changes: 1 addition & 1 deletion native/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -663,7 +663,7 @@ pub extern "system" fn Java_org_apache_datafusion_SessionContext_registerScalarU
let invoke_method = env.get_static_method_id(
&bridge_class_local,
"invokeScalarUdf",
"(Lorg/apache/datafusion/ScalarFunction;JJJJI)V",
"(Lorg/apache/datafusion/ScalarFunction;JJJJ[BJJI)B",
)?;

let java_udf = crate::udf::JavaScalarUdf {
Expand Down
213 changes: 133 additions & 80 deletions native/src/udf.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,18 +19,18 @@

use std::any::Any;
use std::fmt;
use std::sync::Arc;

use datafusion::arrow::array::{make_array, Array, ArrayRef, StructArray};
use datafusion::arrow::datatypes::{DataType, Field, Fields};
use datafusion::arrow::ffi::{from_ffi, to_ffi, FFI_ArrowArray, FFI_ArrowSchema};
use datafusion::common::ScalarValue;
use datafusion::error::DataFusionError;
use datafusion::logical_expr::{
ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, TypeSignature, Volatility,
};
use jni::objects::{GlobalRef, JStaticMethodID, JThrowable};
use jni::signature::{Primitive, ReturnType};
use jni::sys::{jlong, jvalue};
use jni::sys::{jbyte, jlong, jvalue};
use jni::JNIEnv;

pub(crate) struct JavaScalarUdf {
Expand Down Expand Up @@ -99,21 +99,8 @@ impl ScalarUDFImpl for JavaScalarUdf {
) -> datafusion::error::Result<ColumnarValue> {
let number_rows = args.number_rows;

// 1. Materialise scalars to arrays so all columns are length-N.
let arrays: Vec<ArrayRef> = args
.args
.iter()
.map(|cv| cv.clone().into_array(number_rows))
.collect::<datafusion::error::Result<Vec<_>>>()?;

// 2. Build a single struct array carrying all arg columns. Field names/types come
// from the signature's Exact type list (matches what the Java caller declared).
let signature_fields: Vec<Arc<Field>> = match &self.signature.type_signature {
TypeSignature::Exact(types) => types
.iter()
.enumerate()
.map(|(i, ty)| Arc::new(Field::new(format!("arg{}", i), ty.clone(), true)))
.collect(),
let signature_types: &[DataType] = match &self.signature.type_signature {
TypeSignature::Exact(types) => types,
_ => {
return Err(DataFusionError::Internal(
"JavaScalarUdf signature is not Exact; only Signature::exact is supported"
Expand All @@ -122,43 +109,101 @@ impl ScalarUDFImpl for JavaScalarUdf {
}
};

let fields = Fields::from(
signature_fields
.iter()
.map(|f| f.as_ref().clone())
.collect::<Vec<Field>>(),
);
let struct_array = StructArray::try_new_with_length(fields, arrays, None, number_rows)
if args.args.len() != signature_types.len() {
return Err(DataFusionError::Internal(format!(
"Java UDF '{}' called with {} args; signature declares {}",
self.name,
args.args.len(),
signature_types.len()
)));
}

// 1. Partition args by kind. ColumnarValue::Scalar stays as a length-1 array so the Java
// side observes it as a Scalar; ColumnarValue::Array passes through at full length.
let mut array_arrays: Vec<ArrayRef> = Vec::new();
let mut array_fields: Vec<Field> = Vec::new();
let mut scalar_arrays: Vec<ArrayRef> = Vec::new();
let mut scalar_fields: Vec<Field> = Vec::new();
let mut arg_kinds: Vec<u8> = Vec::with_capacity(args.args.len());

for (i, cv) in args.args.iter().enumerate() {
let ty = signature_types[i].clone();
match cv {
ColumnarValue::Array(a) => {
array_fields.push(Field::new(
format!("arg{}", array_arrays.len()),
ty,
true,
));
array_arrays.push(a.clone());
arg_kinds.push(0);
}
ColumnarValue::Scalar(s) => {
let arr = s.to_array_of_size(1)?;
scalar_fields.push(Field::new(
format!("arg{}", scalar_arrays.len()),
ty,
true,
));
scalar_arrays.push(arr);
arg_kinds.push(1);
}
}
}

// 2. Build the two struct arrays. Empty field+array vectors with the appropriate length
// cover nullary and all-one-kind cases.
let array_struct = StructArray::try_new_with_length(
Fields::from(array_fields),
array_arrays,
None,
number_rows,
)
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;
let scalar_struct = StructArray::try_new_with_length(
Fields::from(scalar_fields),
scalar_arrays,
None,
1,
)
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;

let (array_ffi_arr, array_ffi_sch) = to_ffi(&array_struct.into_data())
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;
let (scalar_ffi_arr, scalar_ffi_sch) = to_ffi(&scalar_struct.into_data())
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;
let args_data = struct_array.into_data();
let (args_ffi_array, args_ffi_schema) =
to_ffi(&args_data).map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;

// 3. Pre-allocate empty FFI structs for the result.
let result_ffi_array = FFI_ArrowArray::empty();
let result_ffi_schema = FFI_ArrowSchema::empty();
let result_ffi_arr = FFI_ArrowArray::empty();
let result_ffi_sch = FFI_ArrowSchema::empty();

// 4. Box for stable addresses across the JNI call.
let mut args_array_box = Box::new(args_ffi_array);
let mut args_schema_box = Box::new(args_ffi_schema);
let mut result_array_box = Box::new(result_ffi_array);
let mut result_schema_box = Box::new(result_ffi_schema);

let args_array_addr = args_array_box.as_mut() as *mut _ as jlong;
let args_schema_addr = args_schema_box.as_mut() as *mut _ as jlong;
let result_array_addr = result_array_box.as_mut() as *mut _ as jlong;
let result_schema_addr = result_schema_box.as_mut() as *mut _ as jlong;
let mut array_arr_box = Box::new(array_ffi_arr);
let mut array_sch_box = Box::new(array_ffi_sch);
let mut scalar_arr_box = Box::new(scalar_ffi_arr);
let mut scalar_sch_box = Box::new(scalar_ffi_sch);
let mut result_arr_box = Box::new(result_ffi_arr);
let mut result_sch_box = Box::new(result_ffi_sch);

let array_arr_addr = array_arr_box.as_mut() as *mut _ as jlong;
let array_sch_addr = array_sch_box.as_mut() as *mut _ as jlong;
let scalar_arr_addr = scalar_arr_box.as_mut() as *mut _ as jlong;
let scalar_sch_addr = scalar_sch_box.as_mut() as *mut _ as jlong;
let result_arr_addr = result_arr_box.as_mut() as *mut _ as jlong;
let result_sch_addr = result_sch_box.as_mut() as *mut _ as jlong;

// 5. Attach JNI to current thread.
let mut env = crate::jvm()
.attach_current_thread()
.map_err(|e| DataFusionError::Execution(format!("JNI attach failed: {}", e)))?;

// 6. Call JniBridge.invokeScalarUdf(udf, args*, result*, expectedRowCount).
//
// Build the jvalue argument array for call_static_method_unchecked.
// SAFETY: we build the args inline and pass them immediately; the JObject
// pointed to by udf_global_ref is alive for the duration of this call.
// 6. Build the byte[] for argKinds inside the JVM heap. JNI local; freed when env drops.
let arg_kinds_array = env
.byte_array_from_slice(&arg_kinds)
.map_err(|e| {
DataFusionError::Execution(format!("byte_array_from_slice failed: {}", e))
})?;

let expected_rows = i32::try_from(number_rows).map_err(|_| {
DataFusionError::Execution(format!(
"batch row count {} exceeds i32::MAX; UDFs cannot handle batches larger than 2^31 - 1 rows",
Expand All @@ -167,42 +212,29 @@ impl ScalarUDFImpl for JavaScalarUdf {
})?;

let udf_jobject = self.udf_global_ref.as_obj();
// SAFETY: udf_jobject is derived from a GlobalRef alive for the duration of this
// function. The raw pointer is only read by the JNI call below, which happens
// before any code that could drop udf_global_ref.
let call_args: [jvalue; 6] = [
// ScalarFunction instance
jvalue {
l: udf_jobject.as_raw(),
},
// argsArrayAddr
jvalue { j: args_array_addr },
// argsSchemaAddr
jvalue {
j: args_schema_addr,
},
// resultArrayAddr
jvalue {
j: result_array_addr,
},
// resultSchemaAddr
jvalue {
j: result_schema_addr,
},
// expectedRowCount
// SAFETY: udf_global_ref and arg_kinds_array are alive for the duration of this call.
let call_args: [jvalue; 9] = [
jvalue { l: udf_jobject.as_raw() },
jvalue { j: array_arr_addr },
jvalue { j: array_sch_addr },
jvalue { j: scalar_arr_addr },
jvalue { j: scalar_sch_addr },
jvalue { l: arg_kinds_array.as_raw() },
jvalue { j: result_arr_addr },
jvalue { j: result_sch_addr },
jvalue { i: expected_rows },
];

let call_result = unsafe {
env.call_static_method_unchecked(
&self.bridge_class,
self.invoke_method,
ReturnType::Primitive(Primitive::Void),
ReturnType::Primitive(Primitive::Byte),
&call_args,
)
};

// 7. If Java threw, translate to DataFusionError. Always check exception_check first.
// 7. Java-exception path: translate to DataFusionError.
if env.exception_check().unwrap_or(false) {
let throwable = env.exception_occurred().map_err(|e| {
DataFusionError::Execution(format!("exception_occurred failed: {}", e))
Expand All @@ -211,19 +243,22 @@ impl ScalarUDFImpl for JavaScalarUdf {
let message = jthrowable_to_string(&mut env, &throwable, &self.name);
return Err(DataFusionError::Execution(message));
}
call_result.map_err(|e| DataFusionError::Execution(format!("JNI call failed: {}", e)))?;

// 8. Import result. from_ffi consumes the FFI_ArrowArray.
let result_array = *result_array_box;
let result_schema = *result_schema_box;
// SAFETY: Java's `Data.exportVector` populated `result_array_box` and
// `result_schema_box` in place via the C Data Interface, and the
// exception check above guarantees the call succeeded without
// throwing — so the FFI structs are fully initialized.
let result_data = unsafe { from_ffi(result_array, &result_schema) }

let result_kind: jbyte = call_result
.map_err(|e| DataFusionError::Execution(format!("JNI call failed: {}", e)))?
.b()
.map_err(|e| {
DataFusionError::Execution(format!("invokeScalarUdf return decode failed: {}", e))
})?;

// 8. Import the result vector. from_ffi consumes the FFI_ArrowArray.
let result_array_ffi = *result_arr_box;
let result_schema_ffi = *result_sch_box;
// SAFETY: bridge populated both structs via Arrow C Data Interface; the exception check
// above confirmed no Java exception, so the FFI structs are fully initialised.
let result_data = unsafe { from_ffi(result_array_ffi, &result_schema_ffi) }
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;

// 9. Validate type.
if result_data.data_type() != &self.return_type {
return Err(DataFusionError::Execution(format!(
"Java UDF '{}' returned vector of type {:?}; declared return type was {:?}",
Expand All @@ -234,7 +269,25 @@ impl ScalarUDFImpl for JavaScalarUdf {
}

let array: ArrayRef = make_array(result_data);
Ok(ColumnarValue::Array(array))

match result_kind {
0 => Ok(ColumnarValue::Array(array)),
1 => {
if array.len() != 1 {
return Err(DataFusionError::Internal(format!(
"Java UDF '{}' returned Scalar with length {} (expected 1)",
self.name,
array.len()
)));
}
let scalar = ScalarValue::try_from_array(&array, 0)?;
Ok(ColumnarValue::Scalar(scalar))
}
other => Err(DataFusionError::Internal(format!(
"Java UDF '{}' returned unknown kind byte: {}",
self.name, other
))),
}
}
}

Expand Down