@@ -24,6 +24,7 @@ use std::sync::Arc;
2424use arrow:: array:: RecordBatchReader ;
2525use arrow:: ffi_stream:: ArrowArrayStreamReader ;
2626use arrow:: pyarrow:: FromPyArrow ;
27+ use async_trait:: async_trait;
2728use datafusion:: arrow:: datatypes:: { DataType , Schema , SchemaRef } ;
2829use datafusion:: arrow:: pyarrow:: PyArrowType ;
2930use datafusion:: arrow:: record_batch:: RecordBatch ;
@@ -36,14 +37,16 @@ use datafusion::datasource::listing::{
3637} ;
3738use datafusion:: datasource:: { MemTable , TableProvider } ;
3839use datafusion:: execution:: context:: {
39- DataFilePaths , SQLOptions , SessionConfig , SessionContext , TaskContext ,
40+ DataFilePaths , QueryPlanner , SQLOptions , SessionConfig , SessionContext , TaskContext ,
4041} ;
4142use datafusion:: execution:: disk_manager:: DiskManagerMode ;
4243use datafusion:: execution:: memory_pool:: { FairSpillPool , GreedyMemoryPool , UnboundedMemoryPool } ;
4344use datafusion:: execution:: options:: { ArrowReadOptions , ReadOptions } ;
4445use datafusion:: execution:: runtime_env:: RuntimeEnvBuilder ;
4546use datafusion:: execution:: session_state:: SessionStateBuilder ;
4647use datafusion:: execution:: { FunctionRegistry , TaskContextProvider } ;
48+ use datafusion:: logical_expr:: LogicalPlan ;
49+ use datafusion:: physical_plan:: ExecutionPlan ;
4750use datafusion:: prelude:: {
4851 AvroReadOptions , CsvReadOptions , DataFrame , JsonReadOptions , ParquetReadOptions ,
4952} ;
@@ -53,15 +56,18 @@ use datafusion_ffi::config::extension_options::FFI_ExtensionOptions;
5356use datafusion_ffi:: execution:: FFI_TaskContextProvider ;
5457use datafusion_ffi:: proto:: logical_extension_codec:: FFI_LogicalExtensionCodec ;
5558use datafusion_ffi:: proto:: physical_extension_codec:: FFI_PhysicalExtensionCodec ;
59+ use datafusion_ffi:: query_planner:: FFI_QueryPlanner ;
5660use datafusion_ffi:: table_provider_factory:: FFI_TableProviderFactory ;
5761use datafusion_proto:: logical_plan:: LogicalExtensionCodec ;
5862use datafusion_proto:: physical_plan:: PhysicalExtensionCodec ;
5963use datafusion_python_util:: {
6064 create_logical_extension_capsule, create_physical_extension_capsule,
61- ffi_logical_codec_from_pycapsule, get_global_ctx, get_tokio_runtime,
65+ create_query_planner_capsule, ffi_logical_codec_from_pycapsule,
66+ ffi_query_planner_from_pycapsule, get_global_ctx, get_tokio_runtime,
6267 physical_codec_from_pycapsule, physical_optimizer_rule_from_pycapsule, spawn_future,
6368 wait_for_future,
6469} ;
70+ use datafusion_session:: Session ;
6571use object_store:: ObjectStore ;
6672use pyo3:: IntoPyObjectExt ;
6773use pyo3:: exceptions:: { PyKeyError , PyRuntimeError , PyValueError } ;
@@ -221,6 +227,25 @@ impl PySessionConfig {
221227 }
222228}
223229
230+ #[ derive( Debug , Clone ) ]
231+ struct PythonQueryPlanner {
232+ planner : FFI_QueryPlanner ,
233+ }
234+
235+ #[ async_trait]
236+ impl QueryPlanner for PythonQueryPlanner {
237+ async fn create_physical_plan (
238+ & self ,
239+ logical_plan : & LogicalPlan ,
240+ session : & dyn Session ,
241+ ) -> datafusion:: common:: Result < Arc < dyn ExecutionPlan > > {
242+ let runtime = get_tokio_runtime ( ) . handle ( ) . clone ( ) ;
243+ self . planner
244+ . create_physical_plan_with_session_runtime ( logical_plan, session, Some ( runtime) )
245+ . await
246+ }
247+ }
248+
224249/// Runtime options for a SessionContext
225250#[ pyclass(
226251 from_py_object,
@@ -1211,6 +1236,23 @@ impl PySessionContext {
12111236 Ok ( ( ) )
12121237 }
12131238
1239+ pub fn with_query_planner ( & self , planner : Bound < ' _ , PyAny > ) -> PyDataFusionResult < Self > {
1240+ let mut planner = ffi_query_planner_from_pycapsule ( & planner) ?;
1241+ planner. logical_codec = self . ffi_logical_codec ( ) . as_ref ( ) . clone ( ) ;
1242+ planner. physical_codec = self . ffi_physical_codec ( ) . as_ref ( ) . clone ( ) ;
1243+ let planner = Arc :: new ( PythonQueryPlanner { planner } ) ;
1244+ let state = SessionStateBuilder :: new_from_existing ( self . ctx . state ( ) )
1245+ . with_query_planner ( planner)
1246+ . build ( ) ;
1247+ let ctx = Arc :: new ( SessionContext :: new_with_state ( state) ) ;
1248+
1249+ Ok ( Self {
1250+ ctx,
1251+ logical_codec : Arc :: clone ( & self . logical_codec ) ,
1252+ physical_codec : Arc :: clone ( & self . physical_codec ) ,
1253+ } )
1254+ }
1255+
12141256 pub fn table_provider ( & self , name : & str , py : Python ) -> PyResult < PyTable > {
12151257 let provider = wait_for_future ( py, self . ctx . table_provider ( name) )
12161258 // Outer error: runtime/async failure
@@ -1385,6 +1427,19 @@ impl PySessionContext {
13851427 create_logical_extension_capsule ( py, ffi. as_ref ( ) )
13861428 }
13871429
1430+ pub fn __datafusion_query_planner__ < ' py > (
1431+ & self ,
1432+ py : Python < ' py > ,
1433+ ) -> PyResult < Bound < ' py , PyCapsule > > {
1434+ let planner = Arc :: clone ( self . ctx . state ( ) . query_planner ( ) ) ;
1435+ let ffi = FFI_QueryPlanner :: new_with_ffi_codecs (
1436+ planner,
1437+ self . ffi_logical_codec ( ) . as_ref ( ) . clone ( ) ,
1438+ self . ffi_physical_codec ( ) . as_ref ( ) . clone ( ) ,
1439+ ) ;
1440+ create_query_planner_capsule ( py, & ffi)
1441+ }
1442+
13881443 pub fn with_logical_extension_codec < ' py > (
13891444 & self ,
13901445 codec : Bound < ' py , PyAny > ,
0 commit comments