1use 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 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 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}