Skip to main content

datafusion_python/
catalog.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use std::collections::HashSet;
19use std::ptr::NonNull;
20use std::sync::Arc;
21
22use async_trait::async_trait;
23use datafusion::catalog::{
24    CatalogProvider, CatalogProviderList, MemoryCatalogProvider, MemoryCatalogProviderList,
25    MemorySchemaProvider, SchemaProvider,
26};
27use datafusion::common::DataFusionError;
28use datafusion::datasource::TableProvider;
29use datafusion_ffi::catalog_provider::FFI_CatalogProvider;
30use datafusion_ffi::proto::logical_extension_codec::FFI_LogicalExtensionCodec;
31use datafusion_ffi::schema_provider::FFI_SchemaProvider;
32use datafusion_python_util::{
33    create_logical_extension_capsule, ffi_logical_codec_from_pycapsule, wait_for_future,
34};
35use pyo3::IntoPyObjectExt;
36use pyo3::exceptions::PyKeyError;
37use pyo3::prelude::*;
38use pyo3::types::PyCapsule;
39
40use crate::context::PySessionContext;
41use crate::dataset::Dataset;
42use crate::errors::{PyDataFusionError, PyDataFusionResult, py_datafusion_err, to_datafusion_err};
43use crate::table::PyTable;
44
45#[pyclass(
46    from_py_object,
47    frozen,
48    name = "RawCatalogList",
49    module = "datafusion.catalog",
50    subclass
51)]
52#[derive(Clone)]
53pub struct PyCatalogList {
54    pub catalog_list: Arc<dyn CatalogProviderList>,
55    codec: Arc<FFI_LogicalExtensionCodec>,
56}
57
58#[pyclass(
59    from_py_object,
60    frozen,
61    name = "RawCatalog",
62    module = "datafusion.catalog",
63    subclass
64)]
65#[derive(Clone)]
66pub struct PyCatalog {
67    pub catalog: Arc<dyn CatalogProvider>,
68    codec: Arc<FFI_LogicalExtensionCodec>,
69}
70
71#[pyclass(
72    from_py_object,
73    frozen,
74    name = "RawSchema",
75    module = "datafusion.catalog",
76    subclass
77)]
78#[derive(Clone)]
79pub struct PySchema {
80    pub schema: Arc<dyn SchemaProvider>,
81    codec: Arc<FFI_LogicalExtensionCodec>,
82}
83
84impl PyCatalog {
85    pub(crate) fn new_from_parts(
86        catalog: Arc<dyn CatalogProvider>,
87        codec: Arc<FFI_LogicalExtensionCodec>,
88    ) -> Self {
89        Self { catalog, codec }
90    }
91}
92
93impl PySchema {
94    pub(crate) fn new_from_parts(
95        schema: Arc<dyn SchemaProvider>,
96        codec: Arc<FFI_LogicalExtensionCodec>,
97    ) -> Self {
98        Self { schema, codec }
99    }
100}
101
102#[pymethods]
103impl PyCatalogList {
104    #[new]
105    pub fn new(
106        py: Python,
107        catalog_list: Py<PyAny>,
108        session: Option<Bound<PyAny>>,
109    ) -> PyResult<Self> {
110        let codec = extract_logical_extension_codec(py, session)?;
111        let catalog_list = Arc::new(RustWrappedPyCatalogProviderList::new(
112            catalog_list,
113            codec.clone(),
114        )) as Arc<dyn CatalogProviderList>;
115        Ok(Self {
116            catalog_list,
117            codec,
118        })
119    }
120
121    #[staticmethod]
122    pub fn memory_catalog_list(py: Python, session: Option<Bound<PyAny>>) -> PyResult<Self> {
123        let codec = extract_logical_extension_codec(py, session)?;
124        let catalog_list =
125            Arc::new(MemoryCatalogProviderList::default()) as Arc<dyn CatalogProviderList>;
126        Ok(Self {
127            catalog_list,
128            codec,
129        })
130    }
131
132    pub fn catalog_names(&self) -> HashSet<String> {
133        self.catalog_list.catalog_names().into_iter().collect()
134    }
135
136    #[pyo3(signature = (name="public"))]
137    pub fn catalog(&self, name: &str) -> PyResult<Py<PyAny>> {
138        let catalog = self
139            .catalog_list
140            .catalog(name)
141            .ok_or(PyKeyError::new_err(format!(
142                "Schema with name {name} doesn't exist."
143            )))?;
144
145        Python::attach(
146            |py| match catalog.downcast_ref::<RustWrappedPyCatalogProvider>() {
147                Some(wrapped_catalog) => Ok(wrapped_catalog.catalog_provider.clone_ref(py)),
148                None => PyCatalog::new_from_parts(catalog, self.codec.clone()).into_py_any(py),
149            },
150        )
151    }
152
153    pub fn register_catalog(&self, name: &str, catalog_provider: Bound<'_, PyAny>) -> PyResult<()> {
154        let provider = extract_catalog_provider_from_pyobj(catalog_provider, self.codec.as_ref())?;
155
156        let _ = self
157            .catalog_list
158            .register_catalog(name.to_owned(), provider);
159
160        Ok(())
161    }
162
163    pub fn __repr__(&self) -> PyResult<String> {
164        let mut names: Vec<String> = self.catalog_names().into_iter().collect();
165        names.sort();
166        Ok(format!("CatalogList(catalog_names=[{}])", names.join(", ")))
167    }
168}
169
170#[pymethods]
171impl PyCatalog {
172    #[new]
173    pub fn new(py: Python, catalog: Py<PyAny>, session: Option<Bound<PyAny>>) -> PyResult<Self> {
174        let codec = extract_logical_extension_codec(py, session)?;
175        let catalog = Arc::new(RustWrappedPyCatalogProvider::new(catalog, codec.clone()))
176            as Arc<dyn CatalogProvider>;
177        Ok(Self { catalog, codec })
178    }
179
180    #[staticmethod]
181    pub fn memory_catalog(py: Python, session: Option<Bound<PyAny>>) -> PyResult<Self> {
182        let codec = extract_logical_extension_codec(py, session)?;
183        let catalog = Arc::new(MemoryCatalogProvider::default()) as Arc<dyn CatalogProvider>;
184        Ok(Self { catalog, codec })
185    }
186
187    pub fn schema_names(&self) -> HashSet<String> {
188        self.catalog.schema_names().into_iter().collect()
189    }
190
191    #[pyo3(signature = (name="public"))]
192    pub fn schema(&self, name: &str) -> PyResult<Py<PyAny>> {
193        let schema = self
194            .catalog
195            .schema(name)
196            .ok_or(PyKeyError::new_err(format!(
197                "Schema with name {name} doesn't exist."
198            )))?;
199
200        Python::attach(
201            |py| match schema.downcast_ref::<RustWrappedPySchemaProvider>() {
202                Some(wrapped_schema) => Ok(wrapped_schema.schema_provider.clone_ref(py)),
203                None => PySchema::new_from_parts(schema, self.codec.clone()).into_py_any(py),
204            },
205        )
206    }
207
208    pub fn register_schema(&self, name: &str, schema_provider: Bound<'_, PyAny>) -> PyResult<()> {
209        let provider = extract_schema_provider_from_pyobj(schema_provider, self.codec.as_ref())?;
210
211        let _ = self
212            .catalog
213            .register_schema(name, provider)
214            .map_err(py_datafusion_err)?;
215
216        Ok(())
217    }
218
219    pub fn deregister_schema(&self, name: &str, cascade: bool) -> PyResult<()> {
220        let _ = self
221            .catalog
222            .deregister_schema(name, cascade)
223            .map_err(py_datafusion_err)?;
224
225        Ok(())
226    }
227
228    pub fn __repr__(&self) -> PyResult<String> {
229        let mut names: Vec<String> = self.schema_names().into_iter().collect();
230        names.sort();
231        Ok(format!("Catalog(schema_names=[{}])", names.join(", ")))
232    }
233}
234
235#[pymethods]
236impl PySchema {
237    #[new]
238    pub fn new(
239        py: Python,
240        schema_provider: Py<PyAny>,
241        session: Option<Bound<PyAny>>,
242    ) -> PyResult<Self> {
243        let codec = extract_logical_extension_codec(py, session)?;
244        let schema =
245            Arc::new(RustWrappedPySchemaProvider::new(schema_provider)) as Arc<dyn SchemaProvider>;
246        Ok(Self { schema, codec })
247    }
248
249    #[staticmethod]
250    fn memory_schema(py: Python, session: Option<Bound<PyAny>>) -> PyResult<Self> {
251        let codec = extract_logical_extension_codec(py, session)?;
252        let schema = Arc::new(MemorySchemaProvider::default()) as Arc<dyn SchemaProvider>;
253        Ok(Self { schema, codec })
254    }
255
256    #[getter]
257    fn table_names(&self) -> HashSet<String> {
258        self.schema.table_names().into_iter().collect()
259    }
260
261    fn table(&self, name: &str, py: Python) -> PyDataFusionResult<PyTable> {
262        if let Some(table) = wait_for_future(py, self.schema.table(name))?? {
263            Ok(PyTable::from(table))
264        } else {
265            Err(PyDataFusionError::Common(format!(
266                "Table not found: {name}"
267            )))
268        }
269    }
270
271    fn __repr__(&self) -> PyResult<String> {
272        let mut names: Vec<String> = self.table_names().into_iter().collect();
273        names.sort();
274        Ok(format!("Schema(table_names=[{}])", names.join(";")))
275    }
276
277    fn register_table(&self, name: &str, table_provider: Bound<'_, PyAny>) -> PyResult<()> {
278        let py = table_provider.py();
279        let codec_capsule = create_logical_extension_capsule(py, self.codec.as_ref())?
280            .as_any()
281            .clone();
282
283        let table = PyTable::new(table_provider, Some(codec_capsule))?;
284
285        let _ = self
286            .schema
287            .register_table(name.to_string(), table.table)
288            .map_err(py_datafusion_err)?;
289
290        Ok(())
291    }
292
293    fn deregister_table(&self, name: &str) -> PyResult<()> {
294        let _ = self
295            .schema
296            .deregister_table(name)
297            .map_err(py_datafusion_err)?;
298
299        Ok(())
300    }
301
302    fn table_exist(&self, name: &str) -> bool {
303        self.schema.table_exist(name)
304    }
305}
306
307#[derive(Debug)]
308pub(crate) struct RustWrappedPySchemaProvider {
309    schema_provider: Py<PyAny>,
310    owner_name: Option<String>,
311}
312
313impl RustWrappedPySchemaProvider {
314    pub fn new(schema_provider: Py<PyAny>) -> Self {
315        let owner_name = Python::attach(|py| {
316            schema_provider
317                .bind(py)
318                .getattr("owner_name")
319                .ok()
320                .map(|name| name.to_string())
321        });
322
323        Self {
324            schema_provider,
325            owner_name,
326        }
327    }
328
329    fn table_inner(&self, name: &str) -> PyResult<Option<Arc<dyn TableProvider>>> {
330        Python::attach(|py| {
331            let provider = self.schema_provider.bind(py);
332            let py_table_method = provider.getattr("table")?;
333
334            let py_table = py_table_method.call((name,), None)?;
335            if py_table.is_none() {
336                return Ok(None);
337            }
338
339            let table = PyTable::new(py_table, None)?;
340
341            Ok(Some(table.table))
342        })
343    }
344}
345
346#[async_trait]
347impl SchemaProvider for RustWrappedPySchemaProvider {
348    fn owner_name(&self) -> Option<&str> {
349        self.owner_name.as_deref()
350    }
351
352    fn table_names(&self) -> Vec<String> {
353        Python::attach(|py| {
354            let provider = self.schema_provider.bind(py);
355
356            provider
357                .getattr("table_names")
358                .and_then(|names| names.extract::<Vec<String>>())
359                .unwrap_or_else(|err| {
360                    log::error!("Unable to get table_names: {err}");
361                    Vec::default()
362                })
363        })
364    }
365
366    async fn table(
367        &self,
368        name: &str,
369    ) -> datafusion::common::Result<Option<Arc<dyn TableProvider>>, DataFusionError> {
370        self.table_inner(name)
371            .map_err(|e| DataFusionError::External(Box::new(e)))
372    }
373
374    fn register_table(
375        &self,
376        name: String,
377        table: Arc<dyn TableProvider>,
378    ) -> datafusion::common::Result<Option<Arc<dyn TableProvider>>> {
379        let py_table = PyTable::from(table);
380        Python::attach(|py| {
381            let provider = self.schema_provider.bind(py);
382            let _ = provider
383                .call_method1("register_table", (name, py_table))
384                .map_err(to_datafusion_err)?;
385            // Since the definition of `register_table` says that an error
386            // will be returned if the table already exists, there is no
387            // case where we want to return a table provider as output.
388            Ok(None)
389        })
390    }
391
392    fn deregister_table(
393        &self,
394        name: &str,
395    ) -> datafusion::common::Result<Option<Arc<dyn TableProvider>>> {
396        Python::attach(|py| {
397            let provider = self.schema_provider.bind(py);
398            let table = provider
399                .call_method1("deregister_table", (name,))
400                .map_err(to_datafusion_err)?;
401            if table.is_none() {
402                return Ok(None);
403            }
404
405            // If we can turn this table provider into a `Dataset`, return it.
406            // Otherwise, return None.
407            let dataset = match Dataset::new(&table, py) {
408                Ok(dataset) => Some(Arc::new(dataset) as Arc<dyn TableProvider>),
409                Err(_) => None,
410            };
411
412            Ok(dataset)
413        })
414    }
415
416    fn table_exist(&self, name: &str) -> bool {
417        Python::attach(|py| {
418            let provider = self.schema_provider.bind(py);
419            provider
420                .call_method1("table_exist", (name,))
421                .and_then(|pyobj| pyobj.extract())
422                .unwrap_or(false)
423        })
424    }
425}
426
427#[derive(Debug)]
428pub(crate) struct RustWrappedPyCatalogProvider {
429    pub(crate) catalog_provider: Py<PyAny>,
430    codec: Arc<FFI_LogicalExtensionCodec>,
431}
432
433impl RustWrappedPyCatalogProvider {
434    pub fn new(catalog_provider: Py<PyAny>, codec: Arc<FFI_LogicalExtensionCodec>) -> Self {
435        Self {
436            catalog_provider,
437            codec,
438        }
439    }
440
441    fn schema_inner(&self, name: &str) -> PyResult<Option<Arc<dyn SchemaProvider>>> {
442        Python::attach(|py| {
443            let provider = self.catalog_provider.bind(py);
444
445            let py_schema = provider.call_method1("schema", (name,))?;
446            if py_schema.is_none() {
447                return Ok(None);
448            }
449
450            extract_schema_provider_from_pyobj(py_schema, self.codec.as_ref()).map(Some)
451        })
452    }
453}
454
455#[async_trait]
456impl CatalogProvider for RustWrappedPyCatalogProvider {
457    fn schema_names(&self) -> Vec<String> {
458        Python::attach(|py| {
459            let provider = self.catalog_provider.bind(py);
460            provider
461                .call_method0("schema_names")
462                .and_then(|names| names.extract::<HashSet<String>>())
463                .map(|names| names.into_iter().collect())
464                .unwrap_or_else(|err| {
465                    log::error!("Unable to get schema_names: {err}");
466                    Vec::default()
467                })
468        })
469    }
470
471    fn schema(&self, name: &str) -> Option<Arc<dyn SchemaProvider>> {
472        self.schema_inner(name).unwrap_or_else(|err| {
473            log::error!("CatalogProvider schema returned error: {err}");
474            None
475        })
476    }
477
478    fn register_schema(
479        &self,
480        name: &str,
481        schema: Arc<dyn SchemaProvider>,
482    ) -> datafusion::common::Result<Option<Arc<dyn SchemaProvider>>> {
483        Python::attach(|py| {
484            let py_schema = match schema.downcast_ref::<RustWrappedPySchemaProvider>() {
485                Some(wrapped_schema) => wrapped_schema.schema_provider.as_any(),
486                None => &PySchema::new_from_parts(schema, self.codec.clone())
487                    .into_py_any(py)
488                    .map_err(to_datafusion_err)?,
489            };
490
491            let provider = self.catalog_provider.bind(py);
492            let schema = provider
493                .call_method1("register_schema", (name, py_schema))
494                .map_err(to_datafusion_err)?;
495            if schema.is_none() {
496                return Ok(None);
497            }
498
499            let schema = Arc::new(RustWrappedPySchemaProvider::new(schema.into()))
500                as Arc<dyn SchemaProvider>;
501
502            Ok(Some(schema))
503        })
504    }
505
506    fn deregister_schema(
507        &self,
508        name: &str,
509        cascade: bool,
510    ) -> datafusion::common::Result<Option<Arc<dyn SchemaProvider>>> {
511        Python::attach(|py| {
512            let provider = self.catalog_provider.bind(py);
513            let schema = provider
514                .call_method1("deregister_schema", (name, cascade))
515                .map_err(to_datafusion_err)?;
516            if schema.is_none() {
517                return Ok(None);
518            }
519
520            let schema = Arc::new(RustWrappedPySchemaProvider::new(schema.into()))
521                as Arc<dyn SchemaProvider>;
522
523            Ok(Some(schema))
524        })
525    }
526}
527
528#[derive(Debug)]
529pub(crate) struct RustWrappedPyCatalogProviderList {
530    pub(crate) catalog_provider_list: Py<PyAny>,
531    codec: Arc<FFI_LogicalExtensionCodec>,
532}
533
534impl RustWrappedPyCatalogProviderList {
535    pub fn new(catalog_provider_list: Py<PyAny>, codec: Arc<FFI_LogicalExtensionCodec>) -> Self {
536        Self {
537            catalog_provider_list,
538            codec,
539        }
540    }
541
542    fn catalog_inner(&self, name: &str) -> PyResult<Option<Arc<dyn CatalogProvider>>> {
543        Python::attach(|py| {
544            let provider = self.catalog_provider_list.bind(py);
545
546            let py_schema = provider.call_method1("catalog", (name,))?;
547            if py_schema.is_none() {
548                return Ok(None);
549            }
550
551            extract_catalog_provider_from_pyobj(py_schema, self.codec.as_ref()).map(Some)
552        })
553    }
554}
555
556#[async_trait]
557impl CatalogProviderList for RustWrappedPyCatalogProviderList {
558    fn catalog_names(&self) -> Vec<String> {
559        Python::attach(|py| {
560            let provider = self.catalog_provider_list.bind(py);
561            provider
562                .call_method0("catalog_names")
563                .and_then(|names| names.extract::<HashSet<String>>())
564                .map(|names| names.into_iter().collect())
565                .unwrap_or_else(|err| {
566                    log::error!("Unable to get catalog_names: {err}");
567                    Vec::default()
568                })
569        })
570    }
571
572    fn catalog(&self, name: &str) -> Option<Arc<dyn CatalogProvider>> {
573        self.catalog_inner(name).unwrap_or_else(|err| {
574            log::error!("CatalogProvider catalog returned error: {err}");
575            None
576        })
577    }
578
579    fn register_catalog(
580        &self,
581        name: String,
582        catalog: Arc<dyn CatalogProvider>,
583    ) -> Option<Arc<dyn CatalogProvider>> {
584        Python::attach(|py| {
585            let py_catalog = match catalog.downcast_ref::<RustWrappedPyCatalogProvider>() {
586                Some(wrapped_schema) => wrapped_schema.catalog_provider.as_any().clone_ref(py),
587                None => {
588                    match PyCatalog::new_from_parts(catalog, self.codec.clone()).into_py_any(py) {
589                        Ok(c) => c,
590                        Err(err) => {
591                            log::error!(
592                                "register_catalog returned error during conversion to PyAny: {err}"
593                            );
594                            return None;
595                        }
596                    }
597                }
598            };
599
600            let provider = self.catalog_provider_list.bind(py);
601            let catalog = match provider.call_method1("register_catalog", (name, py_catalog)) {
602                Ok(c) => c,
603                Err(err) => {
604                    log::error!("register_catalog returned error: {err}");
605                    return None;
606                }
607            };
608            if catalog.is_none() {
609                return None;
610            }
611
612            let catalog = Arc::new(RustWrappedPyCatalogProvider::new(
613                catalog.into(),
614                self.codec.clone(),
615            )) as Arc<dyn CatalogProvider>;
616
617            Some(catalog)
618        })
619    }
620}
621
622fn extract_catalog_provider_from_pyobj(
623    mut catalog_provider: Bound<PyAny>,
624    codec: &FFI_LogicalExtensionCodec,
625) -> PyResult<Arc<dyn CatalogProvider>> {
626    if catalog_provider.hasattr("__datafusion_catalog_provider__")? {
627        let py = catalog_provider.py();
628        let codec_capsule = create_logical_extension_capsule(py, codec)?;
629        catalog_provider = catalog_provider
630            .getattr("__datafusion_catalog_provider__")?
631            .call1((codec_capsule,))?;
632    }
633
634    let provider = if let Ok(capsule) = catalog_provider.cast::<PyCapsule>() {
635        let data: NonNull<FFI_CatalogProvider> = capsule
636            .pointer_checked(Some(c"datafusion_catalog_provider"))?
637            .cast();
638        let provider = unsafe { data.as_ref() };
639        let provider: Arc<dyn CatalogProvider> = provider.into();
640        provider
641    } else {
642        match catalog_provider.extract::<PyCatalog>() {
643            Ok(py_catalog) => py_catalog.catalog,
644            Err(_) => Arc::new(RustWrappedPyCatalogProvider::new(
645                catalog_provider.into(),
646                Arc::new(codec.clone()),
647            )) as Arc<dyn CatalogProvider>,
648        }
649    };
650
651    Ok(provider)
652}
653
654fn extract_schema_provider_from_pyobj(
655    mut schema_provider: Bound<PyAny>,
656    codec: &FFI_LogicalExtensionCodec,
657) -> PyResult<Arc<dyn SchemaProvider>> {
658    if schema_provider.hasattr("__datafusion_schema_provider__")? {
659        let py = schema_provider.py();
660        let codec_capsule = create_logical_extension_capsule(py, codec)?;
661        schema_provider = schema_provider
662            .getattr("__datafusion_schema_provider__")?
663            .call1((codec_capsule,))?;
664    }
665
666    let provider = if let Ok(capsule) = schema_provider.cast::<PyCapsule>() {
667        let data: NonNull<FFI_SchemaProvider> = capsule
668            .pointer_checked(Some(c"datafusion_schema_provider"))?
669            .cast();
670        let provider = unsafe { data.as_ref() };
671        let provider: Arc<dyn SchemaProvider> = provider.into();
672        provider
673    } else {
674        match schema_provider.extract::<PySchema>() {
675            Ok(py_schema) => py_schema.schema,
676            Err(_) => Arc::new(RustWrappedPySchemaProvider::new(schema_provider.into()))
677                as Arc<dyn SchemaProvider>,
678        }
679    };
680
681    Ok(provider)
682}
683
684fn extract_logical_extension_codec(
685    py: Python,
686    obj: Option<Bound<PyAny>>,
687) -> PyResult<Arc<FFI_LogicalExtensionCodec>> {
688    let obj = match obj {
689        Some(obj) => obj,
690        None => PySessionContext::global_ctx()?.into_bound_py_any(py)?,
691    };
692    ffi_logical_codec_from_pycapsule(obj).map(Arc::new)
693}
694
695pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
696    m.add_class::<PyCatalog>()?;
697    m.add_class::<PySchema>()?;
698    m.add_class::<PyTable>()?;
699
700    Ok(())
701}