1use crate::{ProviderId, RuntimeClient, RuntimeError, RuntimeResult};
2use async_trait::async_trait;
3use std::collections::BTreeMap;
4use std::sync::Arc;
5
6#[async_trait]
8pub trait RuntimeProviderFactory: Send + Sync {
9 fn provider_id(&self) -> &ProviderId;
10
11 async fn create(&self) -> RuntimeResult<Arc<dyn RuntimeClient>>;
12}
13
14#[derive(Default)]
17pub struct RuntimeClientRegistry {
18 factories: BTreeMap<ProviderId, Arc<dyn RuntimeProviderFactory>>,
19}
20
21impl RuntimeClientRegistry {
22 pub fn new() -> Self {
23 Self::default()
24 }
25
26 pub fn register(&mut self, factory: Arc<dyn RuntimeProviderFactory>) -> RuntimeResult<()> {
27 let provider = factory.provider_id().clone();
28 if self.factories.contains_key(&provider) {
29 return Err(RuntimeError::InvalidRequest(format!(
30 "Runtime provider {provider:?} is already registered"
31 )));
32 }
33 self.factories.insert(provider, factory);
34 Ok(())
35 }
36
37 pub fn contains(&self, provider: &ProviderId) -> bool {
38 self.factories.contains_key(provider)
39 }
40
41 pub async fn connect(&self, provider: &ProviderId) -> RuntimeResult<Arc<dyn RuntimeClient>> {
42 let client = self
43 .factories
44 .get(provider)
45 .ok_or_else(|| {
46 RuntimeError::ProviderUnavailable(format!(
47 "provider {:?} is not registered",
48 provider.as_str()
49 ))
50 })?
51 .create()
52 .await?;
53 let capabilities = client.capabilities().await?;
54 capabilities.validate().map_err(RuntimeError::Protocol)?;
55 if &capabilities.provider_id != provider {
56 return Err(RuntimeError::Protocol(format!(
57 "Runtime provider factory {:?} created client reporting {:?}",
58 provider.as_str(),
59 capabilities.provider_id.as_str()
60 )));
61 }
62 Ok(client)
63 }
64}