Skip to content

Commit 53bbca2

Browse files
committed
fix: support pycapsule methods with arguments
1 parent cc2ec5c commit 53bbca2

2 files changed

Lines changed: 61 additions & 22 deletions

File tree

crates/core/src/context.rs

Lines changed: 13 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -53,14 +53,13 @@ use datafusion_ffi::config::extension_options::FFI_ExtensionOptions;
5353
use datafusion_ffi::execution::FFI_TaskContextProvider;
5454
use datafusion_ffi::proto::logical_extension_codec::FFI_LogicalExtensionCodec;
5555
use datafusion_ffi::proto::physical_extension_codec::FFI_PhysicalExtensionCodec;
56-
use datafusion_ffi::table_provider_factory::FFI_TableProviderFactory;
5756
use datafusion_proto::logical_plan::LogicalExtensionCodec;
5857
use datafusion_proto::physical_plan::PhysicalExtensionCodec;
5958
use datafusion_python_util::{
6059
create_logical_extension_capsule, create_physical_extension_capsule,
6160
ffi_logical_codec_from_pycapsule, get_global_ctx, get_tokio_runtime,
6261
physical_codec_from_pycapsule, physical_optimizer_rule_from_pycapsule, spawn_future,
63-
wait_for_future,
62+
table_provider_factory_from_pycapsule, wait_for_future,
6463
};
6564
use object_store::ObjectStore;
6665
use pyo3::IntoPyObjectExt;
@@ -713,30 +712,22 @@ impl PySessionContext {
713712
pub fn register_table_factory(
714713
&self,
715714
format: &str,
716-
mut factory: Bound<'_, PyAny>,
715+
factory: Bound<'_, PyAny>,
717716
) -> PyDataFusionResult<()> {
718-
if factory.hasattr("__datafusion_table_provider_factory__")? {
717+
let factory: Arc<dyn TableProviderFactory> = if factory
718+
.hasattr("__datafusion_table_provider_factory__")?
719+
|| factory.cast::<PyCapsule>().is_ok()
720+
{
719721
let py = factory.py();
720722
let ffi = self.ffi_logical_codec();
721723
let codec_capsule = create_logical_extension_capsule(py, ffi.as_ref())?;
722-
factory = factory
723-
.getattr("__datafusion_table_provider_factory__")?
724-
.call1((codec_capsule,))?;
725-
}
726-
727-
let factory: Arc<dyn TableProviderFactory> =
728-
if let Ok(capsule) = factory.cast::<PyCapsule>().map_err(py_datafusion_err) {
729-
let data: NonNull<FFI_TableProviderFactory> = capsule
730-
.pointer_checked(Some(c"datafusion_table_provider_factory"))?
731-
.cast();
732-
let factory = unsafe { data.as_ref() };
733-
factory.into()
734-
} else {
735-
Arc::new(RustWrappedPyTableProviderFactory::new(
736-
factory.into(),
737-
self.ffi_logical_codec(),
738-
))
739-
};
724+
table_provider_factory_from_pycapsule(&factory, (codec_capsule,))?
725+
} else {
726+
Arc::new(RustWrappedPyTableProviderFactory::new(
727+
factory.into(),
728+
self.ffi_logical_codec(),
729+
))
730+
};
740731

741732
let st = self.ctx.state_ref();
742733
let mut lock = st.write();

crates/util/src/lib.rs

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ use std::ptr::NonNull;
2020
use std::sync::{Arc, OnceLock};
2121
use std::time::Duration;
2222

23+
use datafusion::catalog::TableProviderFactory;
2324
use datafusion::datasource::TableProvider;
2425
use datafusion::execution::TaskContext;
2526
use datafusion::execution::context::SessionContext;
@@ -30,6 +31,7 @@ use datafusion_ffi::physical_optimizer::FFI_PhysicalOptimizerRule;
3031
use datafusion_ffi::proto::logical_extension_codec::FFI_LogicalExtensionCodec;
3132
use datafusion_ffi::proto::physical_extension_codec::FFI_PhysicalExtensionCodec;
3233
use datafusion_ffi::table_provider::FFI_TableProvider;
34+
use datafusion_ffi::table_provider_factory::FFI_TableProviderFactory;
3335
use datafusion_proto::physical_plan::PhysicalExtensionCodec;
3436
use pyo3::exceptions::{PyImportError, PyTypeError, PyValueError};
3537
use pyo3::prelude::*;
@@ -249,6 +251,44 @@ pub fn create_physical_extension_capsule<'py>(
249251
/// instead.
250252
#[macro_export]
251253
macro_rules! from_pycapsule {
254+
($fn_name:ident, $capsule_name:literal, $ffi_type:ty, $output_type:ty, call_args) => {
255+
pub fn $fn_name<'py, A>(
256+
obj: &$crate::pyo3::Bound<'py, $crate::pyo3::PyAny>,
257+
args: A,
258+
) -> $crate::pyo3::PyResult<std::sync::Arc<$output_type>>
259+
where
260+
A: $crate::pyo3::call::PyCallArgs<'py>,
261+
{
262+
use $crate::pyo3::prelude::*;
263+
use $crate::pyo3::types::PyCapsule;
264+
265+
let mut obj = obj.clone();
266+
if obj.hasattr(concat!("__", $capsule_name, "__"))? {
267+
obj = obj
268+
.getattr(concat!("__", $capsule_name, "__"))?
269+
.call1(args)?;
270+
}
271+
let capsule = obj.cast::<PyCapsule>().map_err(|_| {
272+
$crate::errors::py_datafusion_err(concat!(
273+
"Invalid ",
274+
$capsule_name,
275+
". Does not contain PyCapsule object."
276+
))
277+
})?;
278+
$crate::validate_pycapsule(&capsule, $capsule_name)?;
279+
280+
let expected_name = std::ffi::CString::new($capsule_name)
281+
.expect("capsule name must not contain interior NUL bytes");
282+
let data: std::ptr::NonNull<$ffi_type> = capsule
283+
.pointer_checked(Some(expected_name.as_c_str()))?
284+
.cast();
285+
let output_obj = unsafe { data.as_ref() };
286+
let output_obj: std::sync::Arc<$output_type> = output_obj.into();
287+
288+
Ok(output_obj)
289+
}
290+
};
291+
252292
($fn_name:ident, $capsule_name:literal, $ffi_type:ty, $output_type:ty) => {
253293
pub fn $fn_name(
254294
obj: &$crate::pyo3::Bound<$crate::pyo3::PyAny>,
@@ -340,6 +380,14 @@ from_pycapsule!(
340380
dyn PhysicalOptimizerRule + Send + Sync
341381
);
342382

383+
from_pycapsule!(
384+
table_provider_factory_from_pycapsule,
385+
"datafusion_table_provider_factory",
386+
FFI_TableProviderFactory,
387+
dyn TableProviderFactory,
388+
call_args
389+
);
390+
343391
try_from_pycapsule!(
344392
task_context_from_pycapsule,
345393
"datafusion_task_context_provider",

0 commit comments

Comments
 (0)