Skip to main content

polars_python/
on_startup.rs

1#![allow(unsafe_op_in_unsafe_fn)]
2use std::any::Any;
3use std::sync::OnceLock;
4
5use arrow::array::Array;
6use polars::chunked_array::object::ObjectArray;
7use polars::prelude::file_provider::FileProviderReturn;
8use polars::prelude::*;
9use polars_core::chunked_array::object::builder::ObjectChunkedBuilder;
10use polars_core::chunked_array::object::registry::AnonymousObjectBuilder;
11use polars_core::chunked_array::object::{registry, set_polars_allow_extension};
12use polars_error::PolarsWarning;
13use polars_error::abort::register_polars_abort_mechanism;
14use polars_ffi::version_0::SeriesExport;
15use polars_plan::plans::python_df_to_rust;
16use polars_utils::python_convert_registry::{FromPythonConvertRegistry, PythonConvertRegistry};
17use pyo3::IntoPyObjectExt;
18use pyo3::prelude::*;
19use pyo3::types::PyCFunction;
20
21use crate::Wrap;
22use crate::dataframe::PyDataFrame;
23use crate::lazyframe::PyLazyFrame;
24use crate::map::lazy::call_lambda_with_series;
25use crate::prelude::ObjectValue;
26use crate::py_modules::{pl_df, polars, polars_rs};
27use crate::series::PySeries;
28
29fn python_function_caller_series(
30    s: &[Column],
31    output_dtype: Option<DataType>,
32    lambda: &Py<PyAny>,
33) -> PolarsResult<Column> {
34    Python::attach(|py| call_lambda_with_series(py, s, output_dtype, lambda))
35}
36
37fn python_function_caller_df(df: DataFrame, lambda: &Py<PyAny>) -> PolarsResult<DataFrame> {
38    Python::attach(|py| {
39        let pypolars = polars(py).bind(py);
40
41        // create a PySeries struct/object for Python
42        let pydf = PyDataFrame::new(df);
43        // Wrap this PySeries object in the python side Series wrapper
44        let mut python_df_wrapper = pypolars
45            .getattr("wrap_df")
46            .unwrap()
47            .call1((pydf.clone(),))
48            .unwrap();
49
50        if !python_df_wrapper
51            .getattr("_df")
52            .unwrap()
53            .is_instance(polars_rs(py).getattr(py, "PyDataFrame").unwrap().bind(py))
54            .unwrap()
55        {
56            let pldf = pl_df(py).bind(py);
57            let width = pydf.width();
58            // Don't resize the Vec to avoid calling SeriesExport's Drop impl
59            // The import takes ownership and is responsible for dropping
60            let mut columns: Vec<SeriesExport> = Vec::with_capacity(width);
61            unsafe {
62                pydf._export_columns(columns.as_mut_ptr() as usize);
63            }
64            // Wrap this PyDataFrame object in the python side DataFrame wrapper
65            python_df_wrapper = pldf
66                .getattr("_import_columns")
67                .unwrap()
68                .call1((columns.as_mut_ptr() as usize, width))
69                .unwrap();
70        }
71        // call the lambda and get a python side df wrapper
72        let result_df_wrapper = lambda.call1(py, (python_df_wrapper,))?;
73
74        // unpack the wrapper in a PyDataFrame
75        let py_pydf = result_df_wrapper.getattr(py, "_df").map_err(|_| {
76            let pytype = result_df_wrapper.bind(py).get_type();
77            PolarsError::ComputeError(
78                format!("Expected 'LazyFrame.map' to return a 'DataFrame', got a '{pytype}'",)
79                    .into(),
80            )
81        })?;
82        // Downcast to Rust
83        match py_pydf.extract::<PyDataFrame>(py) {
84            Ok(pydf) => Ok(pydf.df.into_inner()),
85            Err(_) => python_df_to_rust(py, result_df_wrapper.into_bound(py)),
86        }
87    })
88}
89
90fn warning_function(msg: &str, warning: PolarsWarning) {
91    Python::attach(|py| {
92        let Some(warn_fn) = WARN_FUNCTION.get() else {
93            eprintln!("{msg}");
94            return;
95        };
96
97        if let Err(e) = warn_fn
98            .bind(py)
99            .call1((msg, Wrap(warning).into_pyobject(py).unwrap()))
100        {
101            eprintln!("{e}")
102        }
103    });
104}
105
106static POLARS_REGISTRY_INIT_LOCK: OnceLock<()> = OnceLock::new();
107static WARN_FUNCTION: OnceLock<Py<PyAny>> = OnceLock::new();
108
109/// # Safety
110/// Caller must ensure that no other threads read the objects set by this registration.
111pub unsafe fn register_startup_deps(catch_keyboard_interrupt: bool, warn_function: Py<PyAny>) {
112    // TODO: should we throw an error if we try to initialize while already initialized?
113    POLARS_REGISTRY_INIT_LOCK.get_or_init(|| {
114        WARN_FUNCTION.set(warn_function).unwrap();
115        set_polars_allow_extension(true);
116
117        // Stack frames can get really large in debug mode.
118        #[cfg(debug_assertions)]
119        {
120            recursive::set_minimum_stack_size(1024 * 1024);
121            recursive::set_stack_allocation_size(1024 * 1024 * 16);
122        }
123
124        #[cfg(feature = "backtrace_filter")]
125        {
126            use std::path::Path;
127            use color_backtrace::{BacktracePrinter, default_output_stream, default_is_dependency_frame, Frame, ColorScheme};
128            use color_backtrace::termcolor::{ColorSpec, Color};
129
130            let polars_base_path = || {
131                let on_startup = Path::new(file!()).canonicalize().ok()?;
132                let src = on_startup.parent()?;
133                let polars_python = src.parent()?;
134                let crates = polars_python.parent()?;
135                let root = crates.parent()?;
136                Some(root.to_path_buf())
137            };
138
139            let mut btp = BacktracePrinter::default();
140            if let Some(bp) = polars_base_path() {
141                btp = btp.dependency_predicate(Box::new(move |frame: &Frame| -> bool {
142                    if let Some(file) = frame.filename.as_ref().and_then(|f| f.canonicalize().ok()) {
143                        !file.starts_with(&bp)
144                    } else {
145                        default_is_dependency_frame(frame)
146                    }
147                }));
148            }
149
150            let mut color_scheme = ColorScheme::classic();
151            color_scheme.dependency_code = ColorSpec::new();
152            color_scheme.dependency_code.set_dimmed(true);
153            color_scheme.dependency_code = color_scheme.dependency_code_hash.clone();
154            color_scheme.crate_code = ColorSpec::new();
155            color_scheme.crate_code.set_fg(Some(Color::Blue));
156            color_scheme.crate_code_hash = color_scheme.crate_code.clone();
157
158            btp
159                .color_scheme(color_scheme)
160                .install(default_output_stream());
161        }
162
163        // Register object type builder.
164        let object_builder = Box::new(|name: PlSmallStr, capacity: usize| {
165            Box::new(ObjectChunkedBuilder::<ObjectValue>::new(name, capacity))
166                as Box<dyn AnonymousObjectBuilder>
167        });
168
169        let object_converter = Arc::new(|av: AnyValue| {
170            let object = Python::attach(|py| ObjectValue {
171                inner: Wrap(av).into_py_any(py).unwrap(),
172            });
173            Box::new(object) as Box<dyn Any>
174        });
175        let pyobject_converter = Arc::new(|av: AnyValue| {
176            let object = Python::attach(|py| Wrap(av).into_py_any(py).unwrap());
177            Box::new(object) as Box<dyn Any>
178        });
179        fn object_array_getter(arr: &dyn Array, idx: usize) -> Option<AnyValue<'_>> {
180            let arr = arr.as_any().downcast_ref::<ObjectArray<ObjectValue>>().unwrap();
181            arr.get(idx).map(|v| AnyValue::Object(v))
182        }
183        fn with_gil(f: &mut dyn FnMut()) {
184            Python::attach(|_| f())
185        }
186
187        polars_utils::python_convert_registry::register_converters(PythonConvertRegistry {
188            from_py: FromPythonConvertRegistry {
189                file_provider_result: Arc::new(|py_f| {
190                    Python::attach(|py| {
191                        Ok(Box::new(py_f.extract::<Wrap<FileProviderReturn>>(py)?.0) as _)
192                    })
193                }),
194                series: Arc::new(|py_f| {
195                    Python::attach(|py| {
196                        Ok(Box::new(py_f.extract::<PySeries>(py)?.series.into_inner()) as _)
197                    })
198                }),
199                df: Arc::new(|py_f| {
200                    Python::attach(|py| {
201                        Ok(Box::new(py_f.extract::<PyDataFrame>(py)?.df.into_inner()) as _)
202                    })
203                }),
204                dsl_plan: Arc::new(|py_f| {
205                    Python::attach(|py| {
206                        Ok(Box::new(
207                            py_f.extract::<PyLazyFrame>(py)?
208                                .ldf
209                                .into_inner()
210                                .logical_plan,
211                        ) as _)
212                    })
213                }),
214                schema: Arc::new(|py_f| {
215                    Python::attach(|py| {
216                        Ok(Box::new(py_f.extract::<Wrap<polars_core::schema::Schema>>(py)?.0) as _)
217                    })
218                }),
219            },
220            to_py: polars_utils::python_convert_registry::ToPythonConvertRegistry {
221                df: Arc::new(|df| {
222                    Python::attach(|py| {
223                        PyDataFrame::new(df.downcast_ref::<DataFrame>().unwrap().clone())
224                            .into_py_any(py)
225                    })
226                }),
227                series: Arc::new(|series| {
228                    Python::attach(|py| {
229                        PySeries::new(series.downcast_ref::<Series>().unwrap().clone())
230                            .into_py_any(py)
231                    })
232                }),
233                dsl_plan: Arc::new(|dsl_plan| {
234                    Python::attach(|py| {
235                        PyLazyFrame::from(LazyFrame::from(
236                            dsl_plan
237                                .downcast_ref::<polars_plan::dsl::DslPlan>()
238                                .unwrap()
239                                .clone(),
240                        ))
241                        .into_py_any(py)
242                    })
243                }),
244                schema: Arc::new(|schema| {
245                    Python::attach(|py| {
246                        Wrap(
247                            schema
248                                .downcast_ref::<polars_core::schema::Schema>()
249                                .unwrap()
250                                .clone(),
251                        )
252                        .into_py_any(py)
253                    })
254                }),
255            },
256        });
257
258        let object_size = size_of::<ObjectValue>();
259        let physical_dtype = ArrowDataType::FixedSizeBinary(object_size);
260        registry::register_object_builder(
261            object_builder,
262            object_converter,
263            pyobject_converter,
264            physical_dtype,
265            Arc::new(object_array_getter),
266            Arc::new(with_gil)
267        );
268
269        use crate::dataset::dataset_provider_funcs;
270
271        polars_plan::dsl::DATASET_PROVIDER_VTABLE.get_or_init(|| PythonDatasetProviderVTable {
272            name: dataset_provider_funcs::name,
273            schema: dataset_provider_funcs::schema,
274            to_dataset_scan: dataset_provider_funcs::to_dataset_scan,
275        });
276
277        use crate::delta::dv_provider_funcs;
278
279        polars_plan::dsl::deletion::DELTA_DV_PROVIDER_VTABLE.get_or_init(|| {
280            polars_plan::dsl::deletion::DeltaDeletionVectorProviderVTable {
281                call: dv_provider_funcs::call,
282            }
283        });
284
285        // Register SERIES UDF.
286        python_dsl::CALL_PYTHON_COLUMNS_UDF.set(python_function_caller_series).unwrap();
287        // Register DATAFRAME UDF.
288        python_dsl::CALL_PYTHON_DF_UDF.set(python_function_caller_df).unwrap();
289        // Register warning function for `polars_warn!`.
290        polars_error::set_warning_function(warning_function);
291
292        if catch_keyboard_interrupt {
293            register_polars_abort_mechanism();
294        }
295
296        use polars_core::datatypes::extension::UnknownExtensionTypeBehavior;
297        let behavior = match std::env::var("POLARS_UNKNOWN_EXTENSION_TYPE_BEHAVIOR").as_deref() {
298            Ok("load_as_storage") => UnknownExtensionTypeBehavior::LoadAsStorage,
299            Ok("load_as_extension") => UnknownExtensionTypeBehavior::LoadAsGeneric,
300            Ok("") | Err(_) => UnknownExtensionTypeBehavior::WarnAndLoadAsStorage,
301            _ => {
302                polars_warn!("Invalid value for 'POLARS_UNKNOWN_EXTENSION_TYPE_BEHAVIOR' environment variable. Expected one of 'load_as_storage' or 'load_as_extension'.");
303                UnknownExtensionTypeBehavior::WarnAndLoadAsStorage
304            },
305        };
306        polars_core::datatypes::extension::set_unknown_extension_type_behavior(behavior);
307
308        // Out-of-core cleaning.
309        polars_ooc::init_ooc_cleaner();
310        Python::attach(|py| {
311            let atexit = py.import("atexit").unwrap();
312            let flush_fn = PyCFunction::new_closure(py, None, None, |_args, _kwargs| {
313                polars_ooc::flush_ooc_cleanup();
314                PyResult::Ok(())
315            })
316            .unwrap();
317            atexit.call_method1("register", (flush_fn,)).unwrap();
318        });
319
320    });
321}