Skip to main content

a3s_runtime/
registry.rs

1use crate::{ProviderId, RuntimeClient, RuntimeError, RuntimeResult};
2use async_trait::async_trait;
3use std::collections::BTreeMap;
4use std::sync::Arc;
5
6/// Typed construction boundary for one Runtime provider implementation.
7#[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/// Registry of provider factories. Selection policy belongs to the caller;
15/// this registry never falls back to a default provider.
16#[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}