Skip to main content

iceberg_catalog_loader/
lib.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::HashMap;
19use std::sync::Arc;
20
21use async_trait::async_trait;
22use iceberg::encryption::kms::KmsClientFactory;
23use iceberg::io::StorageFactory;
24use iceberg::{Catalog, CatalogBuilder, Error, ErrorKind, Result, Runtime};
25use iceberg_catalog_glue::GlueCatalogBuilder;
26use iceberg_catalog_hms::HmsCatalogBuilder;
27use iceberg_catalog_rest::RestCatalogBuilder;
28use iceberg_catalog_s3tables::S3TablesCatalogBuilder;
29use iceberg_catalog_sql::SqlCatalogBuilder;
30
31/// A CatalogBuilderFactory creating a new catalog builder.
32type CatalogBuilderFactory = fn() -> Box<dyn BoxedCatalogBuilder>;
33
34/// A registry of catalog builders.
35static CATALOG_REGISTRY: &[(&str, CatalogBuilderFactory)] = &[
36    ("rest", || Box::new(RestCatalogBuilder::default())),
37    ("glue", || Box::new(GlueCatalogBuilder::default())),
38    ("s3tables", || Box::new(S3TablesCatalogBuilder::default())),
39    ("hms", || Box::new(HmsCatalogBuilder::default())),
40    ("sql", || Box::new(SqlCatalogBuilder::default())),
41];
42
43/// Return the list of supported catalog types.
44pub fn supported_types() -> Vec<&'static str> {
45    CATALOG_REGISTRY.iter().map(|(k, _)| *k).collect()
46}
47
48#[async_trait]
49pub trait BoxedCatalogBuilder: Send {
50    /// Sets the storage factory used to build the catalog's `FileIO`; see
51    /// [`CatalogBuilder::with_storage_factory`].
52    fn with_storage_factory(
53        self: Box<Self>,
54        storage_factory: Arc<dyn StorageFactory>,
55    ) -> Box<dyn BoxedCatalogBuilder>;
56
57    /// Sets the KMS client factory used to enable table encryption; see
58    /// [`CatalogBuilder::with_kms_client_factory`].
59    fn with_kms_client_factory(
60        self: Box<Self>,
61        kms_client_factory: Arc<dyn KmsClientFactory>,
62    ) -> Box<dyn BoxedCatalogBuilder>;
63
64    /// Sets the runtime the catalog, and the tables it creates, spawn their
65    /// tasks on; see [`CatalogBuilder::with_runtime`].
66    fn with_runtime(self: Box<Self>, runtime: Runtime) -> Box<dyn BoxedCatalogBuilder>;
67
68    async fn load(
69        self: Box<Self>,
70        name: String,
71        props: HashMap<String, String>,
72    ) -> Result<Arc<dyn Catalog>>;
73}
74
75#[async_trait]
76impl<T: CatalogBuilder + 'static> BoxedCatalogBuilder for T {
77    fn with_storage_factory(
78        self: Box<Self>,
79        storage_factory: Arc<dyn StorageFactory>,
80    ) -> Box<dyn BoxedCatalogBuilder> {
81        Box::new(CatalogBuilder::with_storage_factory(*self, storage_factory))
82    }
83
84    fn with_kms_client_factory(
85        self: Box<Self>,
86        kms_client_factory: Arc<dyn KmsClientFactory>,
87    ) -> Box<dyn BoxedCatalogBuilder> {
88        Box::new(CatalogBuilder::with_kms_client_factory(
89            *self,
90            kms_client_factory,
91        ))
92    }
93
94    fn with_runtime(self: Box<Self>, runtime: Runtime) -> Box<dyn BoxedCatalogBuilder> {
95        Box::new(CatalogBuilder::with_runtime(*self, runtime))
96    }
97
98    async fn load(
99        self: Box<Self>,
100        name: String,
101        props: HashMap<String, String>,
102    ) -> Result<Arc<dyn Catalog>> {
103        let builder = *self;
104        Ok(Arc::new(builder.load(name, props).await?) as Arc<dyn Catalog>)
105    }
106}
107
108/// Load a catalog from a string.
109pub fn load(r#type: &str) -> Result<Box<dyn BoxedCatalogBuilder>> {
110    let key = r#type.trim();
111    if let Some((_, factory)) = CATALOG_REGISTRY
112        .iter()
113        .find(|(k, _)| k.eq_ignore_ascii_case(key))
114    {
115        Ok(factory())
116    } else {
117        Err(Error::new(
118            ErrorKind::FeatureUnsupported,
119            format!(
120                "Unsupported catalog type: {}. Supported types: {}",
121                r#type,
122                supported_types().join(", ")
123            ),
124        ))
125    }
126}
127
128/// Ergonomic catalog loader builder pattern.
129pub struct CatalogLoader<'a> {
130    catalog_type: &'a str,
131}
132
133impl<'a> From<&'a str> for CatalogLoader<'a> {
134    fn from(s: &'a str) -> Self {
135        Self { catalog_type: s }
136    }
137}
138
139impl CatalogLoader<'_> {
140    pub async fn load(
141        self,
142        name: String,
143        props: HashMap<String, String>,
144    ) -> Result<Arc<dyn Catalog>> {
145        let builder = load(self.catalog_type)?;
146        builder.load(name, props).await
147    }
148}
149
150#[cfg(test)]
151mod tests {
152    use std::collections::HashMap;
153    use std::sync::{Arc, Mutex};
154
155    use iceberg::encryption::kms::KmsClientFactory;
156    use iceberg::io::{LocalFsStorageFactory, StorageFactory};
157    use iceberg::memory::{MEMORY_CATALOG_WAREHOUSE, MemoryCatalog, MemoryCatalogBuilder};
158    use iceberg::{CatalogBuilder, Result, Runtime};
159    use sqlx::migrate::MigrateDatabase;
160    use tempfile::TempDir;
161
162    use crate::{BoxedCatalogBuilder, CatalogLoader, load};
163
164    #[tokio::test]
165    async fn test_load_unsupported_catalog() {
166        let result = load("unsupported");
167        assert!(result.is_err());
168    }
169
170    #[tokio::test]
171    async fn test_catalog_loader_pattern() {
172        use iceberg_catalog_rest::REST_CATALOG_PROP_URI;
173
174        let catalog = CatalogLoader::from("rest")
175            .load(
176                "rest".to_string(),
177                HashMap::from([
178                    (
179                        REST_CATALOG_PROP_URI.to_string(),
180                        "http://localhost:8080".to_string(),
181                    ),
182                    ("key".to_string(), "value".to_string()),
183                ]),
184            )
185            .await;
186
187        assert!(catalog.is_ok());
188    }
189
190    #[tokio::test]
191    async fn test_catalog_loader_pattern_rest_catalog() {
192        use iceberg_catalog_rest::REST_CATALOG_PROP_URI;
193
194        let catalog_loader = load("rest").unwrap();
195        let catalog = catalog_loader
196            .load(
197                "rest".to_string(),
198                HashMap::from([
199                    (
200                        REST_CATALOG_PROP_URI.to_string(),
201                        "http://localhost:8080".to_string(),
202                    ),
203                    ("key".to_string(), "value".to_string()),
204                ]),
205            )
206            .await;
207
208        assert!(catalog.is_ok());
209    }
210
211    #[tokio::test]
212    async fn test_catalog_loader_pattern_glue_catalog() {
213        use iceberg_catalog_glue::GLUE_CATALOG_PROP_WAREHOUSE;
214
215        let catalog_loader = load("glue").unwrap();
216        let catalog = catalog_loader
217            .load(
218                "glue".to_string(),
219                HashMap::from([
220                    (
221                        GLUE_CATALOG_PROP_WAREHOUSE.to_string(),
222                        "s3://test".to_string(),
223                    ),
224                    ("key".to_string(), "value".to_string()),
225                ]),
226            )
227            .await;
228
229        assert!(catalog.is_ok());
230    }
231
232    #[tokio::test]
233    async fn test_catalog_loader_pattern_s3tables() {
234        use iceberg_catalog_s3tables::S3TABLES_CATALOG_PROP_TABLE_BUCKET_ARN;
235
236        let catalog = CatalogLoader::from("s3tables")
237            .load(
238                "s3tables".to_string(),
239                HashMap::from([
240                    (
241                        S3TABLES_CATALOG_PROP_TABLE_BUCKET_ARN.to_string(),
242                        "arn:aws:s3tables:us-east-1:123456789012:bucket/test".to_string(),
243                    ),
244                    ("key".to_string(), "value".to_string()),
245                ]),
246            )
247            .await;
248
249        assert!(catalog.is_ok());
250    }
251
252    #[tokio::test]
253    async fn test_catalog_loader_pattern_hms_catalog() {
254        use iceberg_catalog_hms::{HMS_CATALOG_PROP_URI, HMS_CATALOG_PROP_WAREHOUSE};
255
256        let catalog_loader = load("hms").unwrap();
257        let catalog = catalog_loader
258            .with_storage_factory(Arc::new(LocalFsStorageFactory))
259            .load(
260                "hms".to_string(),
261                HashMap::from([
262                    (HMS_CATALOG_PROP_URI.to_string(), "127.0.0.1:1".to_string()),
263                    (
264                        HMS_CATALOG_PROP_WAREHOUSE.to_string(),
265                        "s3://warehouse".to_string(),
266                    ),
267                    ("key".to_string(), "value".to_string()),
268                ]),
269            )
270            .await;
271
272        assert!(catalog.is_ok());
273    }
274
275    fn temp_path() -> String {
276        let temp_dir = TempDir::new().unwrap();
277        temp_dir.path().to_str().unwrap().to_string()
278    }
279
280    #[tokio::test]
281    async fn test_catalog_loader_pattern_sql_catalog() {
282        use iceberg_catalog_sql::{SQL_CATALOG_PROP_URI, SQL_CATALOG_PROP_WAREHOUSE};
283
284        let uri = format!("sqlite:{}", temp_path());
285        sqlx::Sqlite::create_database(&uri).await.unwrap();
286
287        let catalog_loader = load("sql").unwrap();
288        let catalog = catalog_loader
289            .with_storage_factory(Arc::new(LocalFsStorageFactory))
290            .load(
291                "sql".to_string(),
292                HashMap::from([
293                    (SQL_CATALOG_PROP_URI.to_string(), uri),
294                    (
295                        SQL_CATALOG_PROP_WAREHOUSE.to_string(),
296                        "s3://warehouse".to_string(),
297                    ),
298                ]),
299            )
300            .await;
301
302        assert!(catalog.is_ok());
303    }
304
305    /// A memory catalog builder that records the runtime it is given.
306    #[derive(Debug, Default)]
307    struct RuntimeRecordingBuilder {
308        runtime: Arc<Mutex<Option<Runtime>>>,
309    }
310
311    impl CatalogBuilder for RuntimeRecordingBuilder {
312        type C = MemoryCatalog;
313
314        fn with_storage_factory(self, _storage_factory: Arc<dyn StorageFactory>) -> Self {
315            self
316        }
317
318        fn with_kms_client_factory(self, _kms_client_factory: Arc<dyn KmsClientFactory>) -> Self {
319            self
320        }
321
322        fn with_runtime(self, runtime: Runtime) -> Self {
323            *self.runtime.lock().unwrap() = Some(runtime);
324            self
325        }
326
327        fn load(
328            self,
329            name: impl Into<String>,
330            props: HashMap<String, String>,
331        ) -> impl Future<Output = Result<MemoryCatalog>> + Send {
332            MemoryCatalogBuilder::default().load(name, props)
333        }
334    }
335
336    #[test]
337    fn test_with_runtime_reaches_the_catalog_builder() {
338        let tokio_runtime = tokio::runtime::Builder::new_multi_thread()
339            .worker_threads(1)
340            .thread_name("loader-test-runtime")
341            .enable_all()
342            .build()
343            .unwrap();
344        let recorded = Arc::new(Mutex::new(None));
345        let builder: Box<dyn BoxedCatalogBuilder> = Box::new(RuntimeRecordingBuilder {
346            runtime: recorded.clone(),
347        });
348
349        tokio_runtime.block_on(async {
350            builder
351                .with_runtime(Runtime::new(&tokio_runtime))
352                .load(
353                    "memory".to_string(),
354                    HashMap::from([(MEMORY_CATALOG_WAREHOUSE.to_string(), temp_path())]),
355                )
356                .await
357                .unwrap();
358
359            // The builder got the runtime passed to the boxed builder: its
360            // tasks run on that runtime's threads.
361            let runtime = recorded.lock().unwrap().take().unwrap();
362            let thread = runtime
363                .io()
364                .spawn(async { std::thread::current().name().map(str::to_string) })
365                .await
366                .unwrap();
367            assert!(
368                thread
369                    .as_deref()
370                    .is_some_and(|n| n.starts_with("loader-test-runtime")),
371                "got: {thread:?}"
372            );
373        });
374    }
375
376    #[tokio::test]
377    async fn test_error_message_includes_supported_types() {
378        let err = match load("does-not-exist") {
379            Ok(_) => panic!("expected error for unsupported type"),
380            Err(e) => e,
381        };
382        let msg = err.message().to_string();
383        assert!(msg.contains("Supported types:"));
384        // Should include at least the built-in type
385        assert!(msg.contains("rest"));
386    }
387}