1use std::collections::HashMap;
15use std::sync::Arc;
16
17use camel_api::datasource::{DatasourceConfig, DatasourceHandle, PoolFactory};
18use camel_api::error::CamelError;
19use camel_api::health::{AsyncHealthCheck, CheckResult};
20use dashmap::DashMap;
21use parking_lot::RwLock;
22use tokio::sync::OnceCell;
23
24use crate::health_registry::HealthCheckRegistry;
25
26type CacheKey = (String, String);
29
30pub struct RuntimeDatasourceCatalog {
33 configs: HashMap<String, DatasourceConfig>,
34 factories: RwLock<HashMap<String, Arc<dyn PoolFactory>>>,
35 pools: DashMap<CacheKey, Arc<OnceCell<DatasourceHandle>>>,
36 health_registry: Option<Arc<HealthCheckRegistry>>,
37}
38
39impl RuntimeDatasourceCatalog {
40 pub fn new(configs: HashMap<String, DatasourceConfig>) -> Self {
41 Self {
42 configs,
43 factories: RwLock::new(HashMap::new()),
44 pools: DashMap::new(),
45 health_registry: None,
46 }
47 }
48
49 pub fn with_health_registry(mut self, registry: Arc<HealthCheckRegistry>) -> Self {
50 self.health_registry = Some(registry);
51 self
52 }
53
54 fn configs(&self) -> &HashMap<String, DatasourceConfig> {
55 &self.configs
56 }
57
58 fn pools(&self) -> &DashMap<CacheKey, Arc<OnceCell<DatasourceHandle>>> {
59 &self.pools
60 }
61
62 fn health_registry(&self) -> Option<&Arc<HealthCheckRegistry>> {
63 self.health_registry.as_ref()
64 }
65
66 fn factories(&self) -> &RwLock<HashMap<String, Arc<dyn PoolFactory>>> {
67 &self.factories
68 }
69}
70
71impl RuntimeDatasourceCatalog {
74 pub(crate) fn get_config(&self, name: &str) -> Option<DatasourceConfig> {
75 self.configs().get(name).cloned()
76 }
77
78 pub(crate) fn register_factory(
79 &self,
80 kind: &str,
81 factory: Arc<dyn PoolFactory>,
82 ) -> Result<(), CamelError> {
83 let mut factories = self.factories().write(); if factories.contains_key(kind) {
85 return Err(CamelError::Config(format!(
86 "factory '{}' already registered",
87 kind
88 )));
89 }
90 factories.insert(kind.to_string(), factory);
91 Ok(())
92 }
93
94 pub(crate) fn get_pool<'a>(
95 &'a self,
96 name: &'a str,
97 ) -> camel_api::datasource::GetPoolFuture<'a> {
98 Box::pin(async move {
99 let config = self.configs().get(name).cloned().ok_or_else(|| {
100 CamelError::Config(format!("datasource '{}' not found in catalog", name))
101 })?;
102
103 let factory = self.resolve_factory(&config)?;
104 let cache_key: CacheKey = (name.to_string(), factory.name().to_string());
105
106 let cell = self
107 .pools()
108 .entry(cache_key)
109 .or_insert_with(|| Arc::new(OnceCell::new()))
110 .clone();
111
112 let handle = cell
113 .get_or_try_init(|| async {
114 let inner = factory.create(&config).await?;
115 let handle =
116 DatasourceHandle::new(name.to_string(), factory.name().to_string(), inner);
117
118 if let Some(registry) = self.health_registry() {
119 let factory_ref = factory.clone();
120 let handle_for_check = handle.clone();
121 let ds_name = name.to_string();
122 registry.register_for_route(
123 &format!("datasource:{}", ds_name),
124 std::sync::Arc::new(DatasourceHealthCheck {
125 check_name: format!("datasource:{}", ds_name),
126 factory: factory_ref,
127 handle: handle_for_check,
128 }),
129 );
130 registry.mark_route_started(&format!("datasource:{}", ds_name));
131 }
132
133 Ok::<DatasourceHandle, CamelError>(handle)
134 })
135 .await?;
136 Ok(handle.clone())
137 })
138 }
139
140 fn resolve_factory(
141 &self,
142 config: &DatasourceConfig,
143 ) -> Result<Arc<dyn PoolFactory>, CamelError> {
144 let factories = self.factories().read(); if let Some(ref provider) = config.provider {
146 let factory = factories.get(provider).ok_or_else(|| {
147 CamelError::Config(format!("unknown datasource provider '{}'", provider))
148 })?;
149 return Ok(factory.clone());
150 }
151
152 let matches: Vec<_> = factories
153 .values()
154 .filter(|entry| entry.matches(config))
155 .collect();
156
157 match matches.len() {
158 0 => Err(CamelError::Config(format!(
159 "no matching factory for datasource url '{}'",
160 scheme_hint(&config.db_url)
161 ))),
162 1 => Ok(matches[0].clone()),
163 _ => {
164 let names: Vec<_> = matches.iter().map(|m| m.name()).collect();
165 Err(CamelError::Config(format!(
166 "ambiguous datasource: {} factories match '{}'. Set explicit 'provider' field.",
167 names.len(),
168 scheme_hint(&config.db_url)
169 )))
170 }
171 }
172 }
173}
174
175fn scheme_hint(db_url: &str) -> String {
178 if let Some(scheme_end) = db_url.find("://") {
179 format!("{}://...", &db_url[..scheme_end])
180 } else {
181 "[REDACTED]".to_string()
182 }
183}
184
185struct DatasourceHealthCheck {
186 check_name: String,
187 factory: Arc<dyn PoolFactory>,
188 handle: DatasourceHandle,
189}
190
191#[async_trait::async_trait]
192impl AsyncHealthCheck for DatasourceHealthCheck {
193 fn name(&self) -> &str {
194 &self.check_name
195 }
196
197 async fn check(&self) -> CheckResult {
198 let status = self.factory.check(&self.handle).await;
199 CheckResult {
200 name: self.check_name.clone(),
201 status,
202 message: None,
203 }
204 }
205}
206
207use camel_api::datasource::DatasourceCatalog;
210
211impl DatasourceCatalog for RuntimeDatasourceCatalog {
212 fn get_config(&self, name: &str) -> Option<DatasourceConfig> {
213 self.get_config(name)
214 }
215
216 fn get_pool<'a>(&'a self, name: &'a str) -> camel_api::datasource::GetPoolFuture<'a> {
217 self.get_pool(name)
218 }
219
220 fn register_factory(
221 &self,
222 kind: &str,
223 factory: Arc<dyn PoolFactory>,
224 ) -> Result<(), CamelError> {
225 self.register_factory(kind, factory)
226 }
227}
228
229#[cfg(test)]
232mod tests {
233 use super::*;
234 use std::any::Any;
235 use std::sync::atomic::{AtomicUsize, Ordering};
236
237 use camel_api::datasource::{CheckFuture, CreatePoolFuture};
238 use camel_api::lifecycle::HealthStatus;
239
240 struct MockFactory {
241 name: &'static str,
242 schemes: &'static [&'static str],
243 create_count: Arc<AtomicUsize>,
244 }
245
246 impl PoolFactory for MockFactory {
247 fn create<'a>(&'a self, _config: &'a DatasourceConfig) -> CreatePoolFuture<'a> {
248 let count = self.create_count.clone();
249 Box::pin(async move {
250 count.fetch_add(1, Ordering::SeqCst);
251 Ok(Arc::new("mock_pool") as Arc<dyn Any + Send + Sync>)
252 })
253 }
254
255 fn check<'a>(&'a self, _handle: &'a DatasourceHandle) -> CheckFuture<'a> {
256 Box::pin(async { HealthStatus::Healthy })
257 }
258
259 fn supported_schemes(&self) -> &[&str] {
260 self.schemes
261 }
262
263 fn name(&self) -> &'static str {
264 self.name
265 }
266 }
267
268 fn make_config(db_url: &str) -> DatasourceConfig {
269 DatasourceConfig {
270 db_url: db_url.to_string(),
271 provider: None,
272 max_connections: None,
273 min_connections: None,
274 idle_timeout_secs: None,
275 max_lifetime_secs: None,
276 ssl_mode: None,
277 ssl_root_cert: None,
278 ssl_cert: None,
279 ssl_key: None,
280 extra: std::collections::HashMap::new(),
281 }
282 }
283
284 #[tokio::test]
285 async fn register_factory_and_get_pool() {
286 let mut configs = HashMap::new();
287 configs.insert(
288 "mydb".to_string(),
289 make_config("postgresql://localhost/mydb"),
290 );
291
292 let catalog = RuntimeDatasourceCatalog::new(configs);
293 let factory = Arc::new(MockFactory {
294 name: "pg",
295 schemes: &["postgresql", "postgres"],
296 create_count: Arc::new(AtomicUsize::new(0)),
297 });
298 catalog.register_factory("postgresql", factory).unwrap();
299
300 let handle = catalog.get_pool("mydb").await.unwrap();
301 assert_eq!(handle.name, "mydb");
302 assert_eq!(handle.provider, "pg");
303 }
304
305 #[tokio::test]
306 async fn shared_pool_for_same_datasource() {
307 let mut configs = HashMap::new();
308 configs.insert(
309 "mydb".to_string(),
310 make_config("postgresql://localhost/mydb"),
311 );
312
313 let count = Arc::new(AtomicUsize::new(0));
314 let catalog = RuntimeDatasourceCatalog::new(configs);
315 let factory = Arc::new(MockFactory {
316 name: "pg",
317 schemes: &["postgresql", "postgres"],
318 create_count: count.clone(),
319 });
320 catalog.register_factory("postgresql", factory).unwrap();
321
322 let h1 = catalog.get_pool("mydb").await.unwrap();
323 let h2 = catalog.get_pool("mydb").await.unwrap();
324
325 assert_eq!(h1.name, h2.name);
326 assert_eq!(h1.provider, h2.provider);
327 assert_eq!(count.load(Ordering::SeqCst), 1);
328 }
329
330 #[tokio::test]
331 async fn unknown_datasource_returns_error() {
332 let configs = HashMap::new();
333 let catalog = RuntimeDatasourceCatalog::new(configs);
334
335 let result = catalog.get_pool("nonexistent").await;
336 assert!(result.is_err());
337 let err = result.unwrap_err();
338 assert!(err.to_string().contains("not found"));
339 }
340
341 #[tokio::test]
342 async fn duplicate_factory_returns_error() {
343 let configs = HashMap::new();
344 let catalog = RuntimeDatasourceCatalog::new(configs);
345 let factory = Arc::new(MockFactory {
346 name: "pg",
347 schemes: &["postgresql"],
348 create_count: Arc::new(AtomicUsize::new(0)),
349 });
350
351 catalog.register_factory("pg", factory.clone()).unwrap();
352 let result = catalog.register_factory("pg", factory);
353 assert!(result.is_err());
354 let err = result.unwrap_err();
355 assert!(err.to_string().contains("already registered"));
356 }
357
358 #[tokio::test]
359 async fn no_matching_factory_returns_error() {
360 let mut configs = HashMap::new();
361 configs.insert("mydb".to_string(), make_config("mongodb://localhost/mydb"));
362
363 let catalog = RuntimeDatasourceCatalog::new(configs);
364 let factory = Arc::new(MockFactory {
365 name: "pg",
366 schemes: &["postgresql"],
367 create_count: Arc::new(AtomicUsize::new(0)),
368 });
369 catalog.register_factory("postgresql", factory).unwrap();
370
371 let result = catalog.get_pool("mydb").await;
372 assert!(result.is_err());
373 let err = result.unwrap_err();
374 assert!(err.to_string().contains("no matching factory"));
375 }
376
377 #[tokio::test]
378 async fn explicit_provider_overrides_scheme() {
379 let mut configs = HashMap::new();
380 configs.insert(
381 "mydb".to_string(),
382 DatasourceConfig {
383 db_url: "postgresql://localhost/mydb".to_string(),
384 provider: Some("mysql_factory".to_string()),
385 max_connections: None,
386 min_connections: None,
387 idle_timeout_secs: None,
388 max_lifetime_secs: None,
389 ssl_mode: None,
390 ssl_root_cert: None,
391 ssl_cert: None,
392 ssl_key: None,
393 extra: std::collections::HashMap::new(),
394 },
395 );
396
397 let pg_count = Arc::new(AtomicUsize::new(0));
398 let mysql_count = Arc::new(AtomicUsize::new(0));
399
400 let catalog = RuntimeDatasourceCatalog::new(configs);
401 let pg_factory = Arc::new(MockFactory {
402 name: "pg",
403 schemes: &["postgresql"],
404 create_count: pg_count.clone(),
405 });
406 let mysql_factory = Arc::new(MockFactory {
407 name: "mysql_factory",
408 schemes: &["mysql"],
409 create_count: mysql_count.clone(),
410 });
411
412 catalog.register_factory("postgresql", pg_factory).unwrap();
413 catalog
414 .register_factory("mysql_factory", mysql_factory)
415 .unwrap();
416
417 let handle = catalog.get_pool("mydb").await.unwrap();
418 assert_eq!(handle.provider, "mysql_factory");
419 assert_eq!(pg_count.load(Ordering::SeqCst), 0);
420 assert_eq!(mysql_count.load(Ordering::SeqCst), 1);
421 }
422
423 #[tokio::test]
424 async fn get_config_returns_clone() {
425 let mut configs = HashMap::new();
426 let original = make_config("postgresql://localhost/mydb");
427 configs.insert("mydb".to_string(), original.clone());
428
429 let catalog = RuntimeDatasourceCatalog::new(configs);
430 let retrieved = catalog.get_config("mydb");
431 assert!(retrieved.is_some());
432 assert_eq!(retrieved.unwrap().db_url, original.db_url);
433 }
434
435 #[tokio::test]
436 async fn get_pool_before_factory_registered_returns_clear_error() {
437 let mut configs = HashMap::new();
438 configs.insert(
439 "mydb".to_string(),
440 make_config("postgresql://localhost/mydb"),
441 );
442
443 let catalog = RuntimeDatasourceCatalog::new(configs);
444
445 let result = catalog.get_pool("mydb").await;
446 assert!(result.is_err());
447 let err = result.unwrap_err();
448 assert!(err.to_string().contains("no matching factory"));
449 }
450
451 #[tokio::test]
452 async fn ambiguous_factory_returns_error() {
453 let mut configs = HashMap::new();
454 configs.insert("orders".into(), make_config("postgres://localhost/test"));
455 let catalog = RuntimeDatasourceCatalog::new(configs);
456 catalog
457 .register_factory(
458 "mock1",
459 Arc::new(MockFactory {
460 name: "mock1",
461 schemes: &["postgres"],
462 create_count: Arc::new(AtomicUsize::new(0)),
463 }),
464 )
465 .unwrap();
466
467 struct MockFactory2;
468 impl PoolFactory for MockFactory2 {
469 fn create<'a>(&'a self, config: &'a DatasourceConfig) -> CreatePoolFuture<'a> {
470 Box::pin(async move {
471 Ok(Arc::new(config.db_url.clone()) as Arc<dyn Any + Send + Sync>)
472 })
473 }
474 fn check<'a>(&'a self, _handle: &'a DatasourceHandle) -> CheckFuture<'a> {
475 Box::pin(async { HealthStatus::Healthy })
476 }
477 fn supported_schemes(&self) -> &[&str] {
478 &["postgres"]
479 }
480 fn name(&self) -> &'static str {
481 "mock2"
482 }
483 }
484 catalog
485 .register_factory("mock2", Arc::new(MockFactory2))
486 .unwrap();
487
488 let result = catalog.get_pool("orders").await;
489 assert!(result.is_err());
490 let msg = result.unwrap_err().to_string();
491 assert!(
492 msg.contains("ambiguous"),
493 "expected ambiguous error, got: {}",
494 msg
495 );
496 }
497
498 #[tokio::test]
499 async fn bad_downcast_returns_clear_error() {
500 let mut configs = HashMap::new();
501 configs.insert(
502 "mydb".to_string(),
503 make_config("postgresql://localhost/mydb"),
504 );
505
506 let catalog = RuntimeDatasourceCatalog::new(configs);
507 let factory = Arc::new(MockFactory {
508 name: "pg",
509 schemes: &["postgresql"],
510 create_count: Arc::new(AtomicUsize::new(0)),
511 });
512 catalog.register_factory("postgresql", factory).unwrap();
513
514 let handle = catalog.get_pool("mydb").await.unwrap();
515
516 let result: Result<Arc<String>, CamelError> = handle.downcast();
517 assert!(result.is_err());
518 let err = result.unwrap_err();
519 assert!(err.to_string().contains("failed to downcast"));
520 assert!(err.to_string().contains("mydb"));
521 assert!(err.to_string().contains("pg"));
522 }
523
524 #[tokio::test]
525 async fn health_check_registered_after_pool_creation() {
526 let mut configs = HashMap::new();
527 configs.insert(
528 "orders".to_string(),
529 make_config("postgresql://localhost/orders"),
530 );
531
532 let registry = Arc::new(crate::health_registry::HealthCheckRegistry::new(
533 std::time::Duration::from_secs(5),
534 ));
535 let catalog = RuntimeDatasourceCatalog::new(configs).with_health_registry(registry.clone());
536 catalog
537 .register_factory(
538 "postgresql",
539 Arc::new(MockFactory {
540 name: "pg",
541 schemes: &["postgresql", "postgres"],
542 create_count: Arc::new(AtomicUsize::new(0)),
543 }),
544 )
545 .unwrap();
546
547 let _ = catalog.get_pool("orders").await.unwrap();
548 registry.mark_route_started("datasource:orders");
549
550 let report = registry.check_all().await;
551 assert!(
552 report
553 .services
554 .iter()
555 .any(|s| s.name.starts_with("datasource:")),
556 "expected datasource health check in report, got: {:?}",
557 report.services
558 );
559 }
560}