Skip to main content

polars_python/dataset/
dataset_provider_funcs.rs

1//! Note: Currently only used for Iceberg / Delta.
2use std::sync::{Arc, LazyLock};
3
4use polars::prelude::{DslPlan, PlSmallStr, Schema, SchemaRef};
5use polars_core::config;
6use polars_error::PolarsResult;
7use polars_plan::plans::PyScanResolveThreadPool;
8use polars_utils::python_function::PythonObject;
9use pyo3::call::PyCallArgs;
10use pyo3::conversion::FromPyObject;
11use pyo3::exceptions::PyValueError;
12use pyo3::pybacked::PyBackedStr;
13use pyo3::types::{PyAnyMethods, PyDict, PyList, PyListMethods};
14use pyo3::{Py, PyAny, PyResult, Python, intern};
15
16use crate::interned;
17use crate::interop::arrow::to_rust::field_to_rust;
18use crate::prelude::{Wrap, get_lf};
19
20pub fn name(dataset_object: &PythonObject) -> PlSmallStr {
21    Python::attach(|py| {
22        PyResult::Ok(PlSmallStr::from_str(
23            &dataset_object
24                .getattr(py, interned::DUNDER_CLASS.get(py))?
25                .getattr(py, interned::DUNDER_NAME.get(py))?
26                .extract::<PyBackedStr>(py)?,
27        ))
28    })
29    .unwrap()
30}
31
32pub fn schema(
33    dataset_object: &PythonObject,
34    py_scan_resolve_threadpool: &PyScanResolveThreadPool,
35) -> PolarsResult<SchemaRef> {
36    Python::attach(|py| {
37        let pyarrow_schema_cls = py
38            .import("pyarrow")
39            .ok()
40            .and_then(|pa| pa.getattr("Schema").ok());
41
42        let schema_obj = py_spawn_call(
43            py,
44            &dataset_object.getattr(py, "schema")?,
45            (),
46            None,
47            py_scan_resolve_threadpool,
48        )?;
49        let schema_cls = schema_obj.getattr(py, interned::DUNDER_CLASS.get(py))?;
50
51        // PyIceberg returns arrow schemas, we convert them here.
52        if let Some(pyarrow_schema_cls) = pyarrow_schema_cls {
53            if schema_cls.is(&pyarrow_schema_cls) {
54                if config::verbose() {
55                    eprintln!("python dataset: convert from arrow schema");
56                }
57
58                let mut iter = schema_obj
59                    .bind(py)
60                    .try_iter()?
61                    .map(|x| x.and_then(field_to_rust));
62
63                let mut last_err = None;
64
65                let schema =
66                    Schema::from_iter_check_duplicates(std::iter::from_fn(|| match iter.next() {
67                        Some(Ok(v)) => Some(v),
68                        Some(Err(e)) => {
69                            last_err = Some(e);
70                            None
71                        },
72                        None => None,
73                    }))?;
74
75                if let Some(last_err) = last_err {
76                    return Err(last_err.into());
77                }
78
79                return Ok(Arc::new(schema));
80            }
81        }
82
83        let Wrap(schema) = Wrap::<Schema>::extract(schema_obj.bind_borrowed(py))?;
84
85        Ok(Arc::new(schema))
86    })
87}
88
89pub fn to_dataset_scan(
90    dataset_object: &PythonObject,
91    existing_resolved_version_key: Option<&str>,
92    limit: Option<usize>,
93    projection: Option<&[PlSmallStr]>,
94    filter_columns: Option<&[PlSmallStr]>,
95    pyarrow_predicate: Option<&str>,
96    py_scan_resolve_threadpool: &PyScanResolveThreadPool,
97) -> PolarsResult<Option<(DslPlan, PlSmallStr)>> {
98    Python::attach(|py| {
99        let kwargs = PyDict::new(py);
100
101        kwargs.set_item(
102            intern!(py, "existing_resolved_version_key"),
103            existing_resolved_version_key,
104        )?;
105
106        if let Some(limit) = limit {
107            kwargs.set_item(intern!(py, "limit"), limit)?;
108        }
109
110        if let Some(projection) = projection {
111            let projection_list = PyList::empty(py);
112
113            for name in projection {
114                projection_list.append(name.as_str())?;
115            }
116
117            kwargs.set_item(intern!(py, "projection"), projection_list)?;
118        }
119
120        if let Some(filter_columns) = filter_columns {
121            let filter_columns_list = PyList::empty(py);
122
123            for name in filter_columns {
124                filter_columns_list.append(name.as_str())?;
125            }
126
127            kwargs.set_item(intern!(py, "filter_columns"), filter_columns_list)?;
128        }
129
130        if let Some(pyarrow_predicate) = pyarrow_predicate {
131            kwargs.set_item(intern!(py, "pyarrow_predicate"), pyarrow_predicate)?;
132        }
133
134        let Some((scan, version)): Option<(Py<PyAny>, Wrap<PlSmallStr>)> = py_spawn_call(
135            py,
136            &dataset_object.getattr(py, intern!(py, "to_dataset_scan"))?,
137            (),
138            Some(&kwargs),
139            py_scan_resolve_threadpool,
140        )?
141        .extract(py)?
142        else {
143            return Ok(None);
144        };
145
146        let Ok(lf) = get_lf(scan.bind(py)) else {
147            return Err(
148                PyValueError::new_err(format!("cannot extract LazyFrame from {}", scan)).into(),
149            );
150        };
151
152        Ok(Some((lf.logical_plan, version.0)))
153    })
154}
155
156fn py_spawn_call<'a>(
157    py: Python<'a>,
158    function: &Py<PyAny>,
159    args: impl PyCallArgs<'a>,
160    kwargs: Option<&pyo3::Bound<'a, PyDict>>,
161    py_scan_resolve_threadpool: &PyScanResolveThreadPool,
162) -> PyResult<Py<PyAny>> {
163    if LazyLock::get(&FN_POOL_WRAP_CLS).is_none() {
164        // Initialization needs GIL, so we must release it to avoid deadlock.
165        py.detach(|| {
166            LazyLock::force(&FN_POOL_WRAP_CLS);
167        })
168    }
169
170    return FN_POOL_WRAP_CLS
171        .call1(py, (function, py_scan_resolve_threadpool))?
172        .call(py, args, kwargs);
173
174    static FN_POOL_WRAP_CLS: LazyLock<Py<PyAny>> = LazyLock::new(|| {
175        Python::attach(|py| {
176            (|| {
177                PyResult::Ok(
178                    py.import("polars._utils.threading")?
179                        .getattr("FnPoolWrap")?
180                        .unbind(),
181                )
182            })()
183            .unwrap()
184        })
185    });
186}