1use parking_lot::RwLock;
19use std::{any::Any, collections::HashMap, fmt, sync::Arc};
20
21#[derive(Clone, Eq, PartialEq, Hash)]
23pub struct TypeKey(&'static str);
24
25impl TypeKey {
26 #[inline]
27 fn of<T: ?Sized + 'static>() -> Self {
28 TypeKey(std::any::type_name::<T>())
29 }
30}
31
32impl fmt::Debug for TypeKey {
33 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
34 f.write_str(self.0)
35 }
36}
37
38#[derive(Clone, Eq, PartialEq, Hash)]
43pub struct ClientScope(Arc<str>);
44
45impl ClientScope {
46 #[inline]
48 #[must_use]
49 pub fn new(scope: impl Into<Arc<str>>) -> Self {
50 Self(scope.into())
51 }
52
53 #[must_use]
57 pub fn gts_id(gts_id: &str) -> Self {
58 let mut s = String::with_capacity("gts:".len() + gts_id.len());
59 s.push_str("gts:");
60 s.push_str(gts_id);
61 Self(Arc::<str>::from(s))
62 }
63
64 #[inline]
65 #[must_use]
66 pub fn as_str(&self) -> &str {
67 &self.0
68 }
69}
70
71impl fmt::Debug for ClientScope {
72 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
73 f.write_str(self.as_str())
74 }
75}
76
77#[derive(Clone, Eq, PartialEq, Hash)]
78struct ScopedKey {
79 type_key: TypeKey,
80 scope: ClientScope,
81}
82
83impl fmt::Debug for ScopedKey {
84 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
85 f.debug_struct("ScopedKey")
86 .field("type_key", &self.type_key)
87 .field("scope", &self.scope)
88 .finish()
89 }
90}
91
92#[derive(Debug, thiserror::Error)]
93pub enum ClientHubError {
94 #[error("client not found: type={type_key:?}")]
95 NotFound { type_key: TypeKey },
96
97 #[error("type mismatch in hub for type={type_key:?}")]
98 TypeMismatch { type_key: TypeKey },
99
100 #[error("scoped client not found: type={type_key:?} scope={scope:?}")]
101 ScopedNotFound {
102 type_key: TypeKey,
103 scope: ClientScope,
104 },
105
106 #[error("type mismatch in hub for type={type_key:?} scope={scope:?}")]
107 ScopedTypeMismatch {
108 type_key: TypeKey,
109 scope: ClientScope,
110 },
111}
112
113type Boxed = Box<dyn Any + Send + Sync>;
114
115type ClientMap = HashMap<TypeKey, Boxed>;
117
118type ScopedClientMap = HashMap<ScopedKey, Boxed>;
120
121#[derive(Default)]
123pub struct ClientHub {
124 map: RwLock<ClientMap>,
125 scoped_map: RwLock<ScopedClientMap>,
126 remote_proxies: RwLock<std::collections::HashSet<TypeKey>>,
135}
136
137impl ClientHub {
138 #[inline]
139 #[must_use]
140 pub fn new() -> Self {
141 Self {
142 map: RwLock::new(HashMap::new()),
143 scoped_map: RwLock::new(HashMap::new()),
144 remote_proxies: RwLock::new(std::collections::HashSet::new()),
145 }
146 }
147}
148
149impl ClientHub {
150 pub fn register<T>(&self, client: Arc<T>)
153 where
154 T: ?Sized + Send + Sync + 'static,
155 {
156 let type_key = TypeKey::of::<T>();
157 let mut w = self.map.write();
158 w.insert(type_key, Box::new(client));
159 }
160
161 pub fn register_scoped<T>(&self, scope: ClientScope, client: Arc<T>)
166 where
167 T: ?Sized + Send + Sync + 'static,
168 {
169 let key = ScopedKey {
170 type_key: TypeKey::of::<T>(),
171 scope,
172 };
173 let mut w = self.scoped_map.write();
174 w.insert(key, Box::new(client));
175 }
176
177 pub fn get<T>(&self) -> Result<Arc<T>, ClientHubError>
183 where
184 T: ?Sized + Send + Sync + 'static,
185 {
186 let type_key = TypeKey::of::<T>();
187 let r = self.map.read();
188
189 let boxed = r.get(&type_key).ok_or(ClientHubError::NotFound {
190 type_key: type_key.clone(),
191 })?;
192
193 if let Some(arc_t) = boxed.downcast_ref::<Arc<T>>() {
195 return Ok(arc_t.clone());
196 }
197 Err(ClientHubError::TypeMismatch { type_key })
198 }
199
200 pub fn try_get<T>(&self) -> Option<Arc<T>>
207 where
208 T: ?Sized + Send + Sync + 'static,
209 {
210 let type_key = TypeKey::of::<T>();
211 let r = self.map.read();
212 let boxed = r.get(&type_key)?;
213 boxed.downcast_ref::<Arc<T>>().cloned()
214 }
215
216 pub fn register_remote_proxy<T>(&self, client: Arc<T>)
223 where
224 T: ?Sized + Send + Sync + 'static,
225 {
226 let type_key = TypeKey::of::<T>();
227 self.remote_proxies.write().insert(type_key.clone());
228 self.map.write().insert(type_key, Box::new(client));
229 }
230
231 pub fn try_get_local<T>(&self) -> Option<Arc<T>>
244 where
245 T: ?Sized + Send + Sync + 'static,
246 {
247 let type_key = TypeKey::of::<T>();
248 if self.remote_proxies.read().contains(&type_key) {
249 return None;
250 }
251 let r = self.map.read();
252 let boxed = r.get(&type_key)?;
253 boxed.downcast_ref::<Arc<T>>().cloned()
254 }
255
256 #[must_use]
261 pub fn has_remote_proxy<T>(&self) -> bool
262 where
263 T: ?Sized + Send + Sync + 'static,
264 {
265 self.remote_proxies.read().contains(&TypeKey::of::<T>())
266 }
267
268 pub fn get_scoped<T>(&self, scope: &ClientScope) -> Result<Arc<T>, ClientHubError>
274 where
275 T: ?Sized + Send + Sync + 'static,
276 {
277 let key = ScopedKey {
278 type_key: TypeKey::of::<T>(),
279 scope: scope.clone(),
280 };
281 let r = self.scoped_map.read();
282
283 let boxed = r.get(&key).ok_or_else(|| ClientHubError::ScopedNotFound {
284 type_key: key.type_key.clone(),
285 scope: key.scope.clone(),
286 })?;
287
288 if let Some(arc_t) = boxed.downcast_ref::<Arc<T>>() {
289 return Ok(arc_t.clone());
290 }
291 Err(ClientHubError::ScopedTypeMismatch {
292 type_key: key.type_key,
293 scope: key.scope,
294 })
295 }
296
297 pub fn try_get_scoped<T>(&self, scope: &ClientScope) -> Option<Arc<T>>
301 where
302 T: ?Sized + Send + Sync + 'static,
303 {
304 let key = ScopedKey {
305 type_key: TypeKey::of::<T>(),
306 scope: scope.clone(),
307 };
308 let r = self.scoped_map.read();
309 let boxed = r.get(&key)?;
310
311 boxed.downcast_ref::<Arc<T>>().cloned()
312 }
313
314 pub fn remove<T>(&self) -> Option<Arc<T>>
316 where
317 T: ?Sized + Send + Sync + 'static,
318 {
319 let type_key = TypeKey::of::<T>();
320 let mut w = self.map.write();
321 let boxed = w.remove(&type_key)?;
322 boxed.downcast::<Arc<T>>().ok().map(|b| *b)
323 }
324
325 pub fn remove_scoped<T>(&self, scope: &ClientScope) -> Option<Arc<T>>
327 where
328 T: ?Sized + Send + Sync + 'static,
329 {
330 let key = ScopedKey {
331 type_key: TypeKey::of::<T>(),
332 scope: scope.clone(),
333 };
334 let mut w = self.scoped_map.write();
335 let boxed = w.remove(&key)?;
336 boxed.downcast::<Arc<T>>().ok().map(|b| *b)
337 }
338
339 pub fn clear(&self) {
341 self.map.write().clear();
342 self.scoped_map.write().clear();
343 }
344
345 pub fn len(&self) -> usize {
347 self.map.read().len() + self.scoped_map.read().len()
348 }
349
350 pub fn is_empty(&self) -> bool {
352 self.map.read().is_empty() && self.scoped_map.read().is_empty()
353 }
354}
355
356#[cfg(test)]
357#[cfg_attr(coverage_nightly, coverage(off))]
358mod tests {
359 use super::*;
360 use toolkit_gts::gts_id;
361
362 #[async_trait::async_trait]
363 trait TestApi: Send + Sync {
364 async fn id(&self) -> usize;
365 }
366
367 struct ImplA(usize);
368 #[async_trait::async_trait]
369 impl TestApi for ImplA {
370 async fn id(&self) -> usize {
371 self.0
372 }
373 }
374
375 #[tokio::test]
376 async fn register_and_get_dyn_trait() {
377 let hub = ClientHub::new();
378 let api: Arc<dyn TestApi> = Arc::new(ImplA(7));
379 hub.register::<dyn TestApi>(api.clone());
380
381 let got = hub.get::<dyn TestApi>().unwrap();
382 assert_eq!(got.id().await, 7);
383 assert_eq!(Arc::as_ptr(&api), Arc::as_ptr(&got));
384 }
385
386 #[tokio::test]
387 async fn remove_works() {
388 let hub = ClientHub::new();
389 let api: Arc<dyn TestApi> = Arc::new(ImplA(42));
390 hub.register::<dyn TestApi>(api);
391
392 assert!(hub.get::<dyn TestApi>().is_ok());
393
394 let removed = hub.remove::<dyn TestApi>();
395 assert!(removed.is_some());
396 assert!(hub.get::<dyn TestApi>().is_err());
397 }
398
399 #[tokio::test]
400 async fn try_get_local_ignores_a_remote_proxy() {
401 let hub = ClientHub::new();
402 let proxy: Arc<dyn TestApi> = Arc::new(ImplA(1));
403 hub.register_remote_proxy::<dyn TestApi>(proxy);
404
405 assert!(hub.try_get::<dyn TestApi>().is_some());
407 assert!(hub.get::<dyn TestApi>().is_ok());
408 assert!(hub.try_get_local::<dyn TestApi>().is_none());
413 assert!(hub.has_remote_proxy::<dyn TestApi>());
414 }
415
416 #[tokio::test]
417 async fn try_get_local_sees_a_plain_registration() {
418 let hub = ClientHub::new();
419 let local: Arc<dyn TestApi> = Arc::new(ImplA(2));
420 hub.register::<dyn TestApi>(local);
421
422 assert!(hub.try_get_local::<dyn TestApi>().is_some());
423 assert!(!hub.has_remote_proxy::<dyn TestApi>());
424 }
425
426 #[tokio::test]
427 async fn overwrite_replaces_atomically() {
428 let hub = ClientHub::new();
429 hub.register::<dyn TestApi>(Arc::new(ImplA(1)));
430
431 let old = hub.get::<dyn TestApi>().unwrap();
432 assert_eq!(old.id().await, 1);
433
434 hub.register::<dyn TestApi>(Arc::new(ImplA(2)));
435
436 let new = hub.get::<dyn TestApi>().unwrap();
437 assert_eq!(new.id().await, 2);
438
439 assert_eq!(old.id().await, 1);
441 }
442
443 #[tokio::test]
444 async fn scoped_register_and_get_dyn_trait() {
445 let hub = ClientHub::new();
446 let scope_a = ClientScope::gts_id(gts_id!(
447 "cf.core.toolkit.plugins.v1~cf.core.tenant_resolver.plugin.v1~contoso.app._.plugin.v1.0"
448 ));
449 let scope_b = ClientScope::gts_id(gts_id!(
450 "cf.core.toolkit.plugins.v1~cf.core.tenant_resolver.plugin.v1~fabrikam.app._.plugin.v1.0"
451 ));
452
453 let api_a: Arc<dyn TestApi> = Arc::new(ImplA(1));
454 let api_b: Arc<dyn TestApi> = Arc::new(ImplA(2));
455
456 hub.register_scoped::<dyn TestApi>(scope_a.clone(), api_a.clone());
457 hub.register_scoped::<dyn TestApi>(scope_b.clone(), api_b.clone());
458
459 assert_eq!(
460 hub.get_scoped::<dyn TestApi>(&scope_a).unwrap().id().await,
461 1
462 );
463 assert_eq!(
464 hub.get_scoped::<dyn TestApi>(&scope_b).unwrap().id().await,
465 2
466 );
467 }
468
469 #[test]
470 fn scoped_get_is_independent_from_global_get() {
471 let hub = ClientHub::new();
472 let scope = ClientScope::gts_id(gts_id!(
473 "cf.core.toolkit.plugins.v1~cf.core.tenant_resolver.plugin.v1~fabrikam.app._.plugin.v1.0"
474 ));
475 hub.register::<str>(Arc::from("global"));
476 hub.register_scoped::<str>(scope.clone(), Arc::from("scoped"));
477
478 assert_eq!(&*hub.get::<str>().unwrap(), "global");
479 assert_eq!(&*hub.get_scoped::<str>(&scope).unwrap(), "scoped");
480 }
481
482 #[test]
483 fn try_get_scoped_returns_some_on_hit() {
484 let hub = ClientHub::new();
485 let scope = ClientScope::gts_id(gts_id!(
486 "cf.core.toolkit.plugins.v1~cf.core.tenant_resolver.plugin.v1~contoso.app._.plugin.v1.0"
487 ));
488 hub.register_scoped::<str>(scope.clone(), Arc::from("scoped"));
489
490 let got = hub.try_get_scoped::<str>(&scope);
491 assert_eq!(got.as_deref(), Some("scoped"));
492 }
493
494 #[test]
495 fn try_get_scoped_returns_none_on_miss() {
496 let hub = ClientHub::new();
497 let scope = ClientScope::gts_id(gts_id!(
498 "cf.core.toolkit.plugins.v1~cf.core.tenant_resolver.plugin.v1~fabrikam.app._.plugin.v1.0"
499 ));
500
501 let got = hub.try_get_scoped::<str>(&scope);
502 assert!(got.is_none());
503 }
504}