1use 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
31type CatalogBuilderFactory = fn() -> Box<dyn BoxedCatalogBuilder>;
33
34static 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
43pub fn supported_types() -> Vec<&'static str> {
45 CATALOG_REGISTRY.iter().map(|(k, _)| *k).collect()
46}
47
48#[async_trait]
49pub trait BoxedCatalogBuilder: Send {
50 fn with_storage_factory(
53 self: Box<Self>,
54 storage_factory: Arc<dyn StorageFactory>,
55 ) -> Box<dyn BoxedCatalogBuilder>;
56
57 fn with_kms_client_factory(
60 self: Box<Self>,
61 kms_client_factory: Arc<dyn KmsClientFactory>,
62 ) -> Box<dyn BoxedCatalogBuilder>;
63
64 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
108pub 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
128pub 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 #[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 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 assert!(msg.contains("rest"));
386 }
387}