Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 13 additions & 22 deletions crates/core/src/context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -53,14 +53,13 @@ use datafusion_ffi::config::extension_options::FFI_ExtensionOptions;
use datafusion_ffi::execution::FFI_TaskContextProvider;
use datafusion_ffi::proto::logical_extension_codec::FFI_LogicalExtensionCodec;
use datafusion_ffi::proto::physical_extension_codec::FFI_PhysicalExtensionCodec;
use datafusion_ffi::table_provider_factory::FFI_TableProviderFactory;
use datafusion_proto::logical_plan::LogicalExtensionCodec;
use datafusion_proto::physical_plan::PhysicalExtensionCodec;
use datafusion_python_util::{
create_logical_extension_capsule, create_physical_extension_capsule,
ffi_logical_codec_from_pycapsule, get_global_ctx, get_tokio_runtime,
physical_codec_from_pycapsule, physical_optimizer_rule_from_pycapsule, spawn_future,
wait_for_future,
table_provider_factory_from_pycapsule, wait_for_future,
};
use object_store::ObjectStore;
use pyo3::IntoPyObjectExt;
Expand Down Expand Up @@ -713,30 +712,22 @@ impl PySessionContext {
pub fn register_table_factory(
&self,
format: &str,
mut factory: Bound<'_, PyAny>,
factory: Bound<'_, PyAny>,
) -> PyDataFusionResult<()> {
if factory.hasattr("__datafusion_table_provider_factory__")? {
let factory: Arc<dyn TableProviderFactory> = if factory
.hasattr("__datafusion_table_provider_factory__")?
|| factory.cast::<PyCapsule>().is_ok()
{
let py = factory.py();
let ffi = self.ffi_logical_codec();
let codec_capsule = create_logical_extension_capsule(py, ffi.as_ref())?;
factory = factory
.getattr("__datafusion_table_provider_factory__")?
.call1((codec_capsule,))?;
}

let factory: Arc<dyn TableProviderFactory> =
if let Ok(capsule) = factory.cast::<PyCapsule>().map_err(py_datafusion_err) {
let data: NonNull<FFI_TableProviderFactory> = capsule
.pointer_checked(Some(c"datafusion_table_provider_factory"))?
.cast();
let factory = unsafe { data.as_ref() };
factory.into()
} else {
Arc::new(RustWrappedPyTableProviderFactory::new(
factory.into(),
self.ffi_logical_codec(),
))
};
table_provider_factory_from_pycapsule(&factory, (codec_capsule,))?
} else {
Arc::new(RustWrappedPyTableProviderFactory::new(
factory.into(),
self.ffi_logical_codec(),
))
};

let st = self.ctx.state_ref();
let mut lock = st.write();
Expand Down
48 changes: 48 additions & 0 deletions crates/util/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ use std::ptr::NonNull;
use std::sync::{Arc, OnceLock};
use std::time::Duration;

use datafusion::catalog::TableProviderFactory;
use datafusion::datasource::TableProvider;
use datafusion::execution::TaskContext;
use datafusion::execution::context::SessionContext;
Expand All @@ -30,6 +31,7 @@ use datafusion_ffi::physical_optimizer::FFI_PhysicalOptimizerRule;
use datafusion_ffi::proto::logical_extension_codec::FFI_LogicalExtensionCodec;
use datafusion_ffi::proto::physical_extension_codec::FFI_PhysicalExtensionCodec;
use datafusion_ffi::table_provider::FFI_TableProvider;
use datafusion_ffi::table_provider_factory::FFI_TableProviderFactory;
use datafusion_proto::physical_plan::PhysicalExtensionCodec;
use pyo3::exceptions::{PyImportError, PyTypeError, PyValueError};
use pyo3::prelude::*;
Expand Down Expand Up @@ -249,6 +251,44 @@ pub fn create_physical_extension_capsule<'py>(
/// instead.
#[macro_export]
macro_rules! from_pycapsule {
($fn_name:ident, $capsule_name:literal, $ffi_type:ty, $output_type:ty, call_args) => {
pub fn $fn_name<'py, A>(
obj: &$crate::pyo3::Bound<'py, $crate::pyo3::PyAny>,
args: A,
) -> $crate::pyo3::PyResult<std::sync::Arc<$output_type>>
where
A: $crate::pyo3::call::PyCallArgs<'py>,
{
use $crate::pyo3::prelude::*;
use $crate::pyo3::types::PyCapsule;

let mut obj = obj.clone();
if obj.hasattr(concat!("__", $capsule_name, "__"))? {
obj = obj
.getattr(concat!("__", $capsule_name, "__"))?
.call1(args)?;
}
let capsule = obj.cast::<PyCapsule>().map_err(|_| {
$crate::errors::py_datafusion_err(concat!(
"Invalid ",
$capsule_name,
". Does not contain PyCapsule object."
))
})?;
$crate::validate_pycapsule(&capsule, $capsule_name)?;

let expected_name = std::ffi::CString::new($capsule_name)
.expect("capsule name must not contain interior NUL bytes");
let data: std::ptr::NonNull<$ffi_type> = capsule
.pointer_checked(Some(expected_name.as_c_str()))?
.cast();
let output_obj = unsafe { data.as_ref() };
let output_obj: std::sync::Arc<$output_type> = output_obj.into();

Ok(output_obj)
}
};

($fn_name:ident, $capsule_name:literal, $ffi_type:ty, $output_type:ty) => {
pub fn $fn_name(
obj: &$crate::pyo3::Bound<$crate::pyo3::PyAny>,
Expand Down Expand Up @@ -340,6 +380,14 @@ from_pycapsule!(
dyn PhysicalOptimizerRule + Send + Sync
);

from_pycapsule!(
table_provider_factory_from_pycapsule,
"datafusion_table_provider_factory",
FFI_TableProviderFactory,
dyn TableProviderFactory,
call_args
);

try_from_pycapsule!(
task_context_from_pycapsule,
"datafusion_task_context_provider",
Expand Down
Loading