Skip to main content

locus_sdk/infrastructure/
registry.rs

1use std::collections::HashMap;
2
3use anyhow::{Result, anyhow};
4
5use crate::domain::ai::{AiCapability, AiProvider, AiProviderRegistry, AiTask, ProviderPolicy};
6
7pub struct InMemoryAiProviderRegistry {
8    providers: HashMap<String, Box<dyn AiProvider>>,
9}
10
11impl InMemoryAiProviderRegistry {
12    pub fn new() -> Self {
13        Self {
14            providers: HashMap::new(),
15        }
16    }
17
18    pub fn register<P>(&mut self, provider: P)
19    where
20        P: AiProvider + 'static,
21    {
22        self.providers
23            .insert(provider.provider_id().to_string(), Box::new(provider));
24    }
25
26    fn provider_supports_task(provider: &dyn AiProvider, task: AiTask) -> bool {
27        let caps = provider.capabilities();
28        match task {
29            AiTask::SemanticEmbedding => caps.contains(&AiCapability::SemanticEmbedding),
30            AiTask::AvecEmbedding => caps.contains(&AiCapability::AvecEmbedding),
31            AiTask::AvecScoring => caps.contains(&AiCapability::AvecScoring),
32        }
33    }
34}
35
36impl Default for InMemoryAiProviderRegistry {
37    fn default() -> Self {
38        Self::new()
39    }
40}
41
42impl AiProviderRegistry for InMemoryAiProviderRegistry {
43    fn resolve(
44        &self,
45        task: AiTask,
46        provider_id: Option<&str>,
47        policy: ProviderPolicy,
48    ) -> Result<&dyn AiProvider> {
49        if let Some(id) = provider_id {
50            let provider = self
51                .providers
52                .get(id)
53                .ok_or_else(|| anyhow!("requested provider '{id}' is not registered"))?;
54            if !Self::provider_supports_task(provider.as_ref(), task) {
55                return Err(anyhow!("provider '{id}' does not support task '{task:?}'"));
56            }
57            return Ok(provider.as_ref());
58        }
59
60        match policy {
61            ProviderPolicy::Required => Err(anyhow!(
62                "provider_policy=required needs an explicit provider_id"
63            )),
64            ProviderPolicy::Auto | ProviderPolicy::Preferred => self
65                .providers
66                .values()
67                .find(|provider| Self::provider_supports_task(provider.as_ref(), task))
68                .map(|provider| provider.as_ref())
69                .ok_or_else(|| anyhow!("no registered provider supports task '{task:?}'")),
70        }
71    }
72
73    fn list_capabilities(&self) -> Vec<(String, Vec<AiCapability>)> {
74        self.providers
75            .iter()
76            .map(|(id, provider)| (id.clone(), provider.capabilities().to_vec()))
77            .collect()
78    }
79}
80
81#[cfg(test)]
82mod tests {
83    use anyhow::Result;
84    use async_trait::async_trait;
85    use locus_core_rs::domain::models::AvecState;
86
87    use super::InMemoryAiProviderRegistry;
88    use crate::domain::ai::{
89        AiCapability, AiProvider, AiProviderRegistry, AiTask, EmbedRequest, ProviderPolicy,
90        ScoreAvecRequest,
91    };
92
93    struct SemanticOnlyProvider;
94
95    #[async_trait]
96    impl AiProvider for SemanticOnlyProvider {
97        fn provider_id(&self) -> &str {
98            "semantic-only"
99        }
100
101        fn capabilities(&self) -> &'static [AiCapability] {
102            &[AiCapability::SemanticEmbedding]
103        }
104
105        async fn embed_semantic(&self, _request: &EmbedRequest) -> Result<Vec<f32>> {
106            Ok(vec![0.1, 0.2, 0.3])
107        }
108
109        async fn embed_avec(&self, _request: &EmbedRequest) -> Result<Vec<f32>> {
110            Ok(vec![0.4, 0.5, 0.6])
111        }
112
113        async fn score_avec(&self, _request: &ScoreAvecRequest) -> Result<AvecState> {
114            Ok(AvecState {
115                stability: 0.5,
116                friction: 0.5,
117                logic: 0.5,
118                autonomy: 0.5,
119            })
120        }
121    }
122
123    struct AvecProvider;
124
125    #[async_trait]
126    impl AiProvider for AvecProvider {
127        fn provider_id(&self) -> &str {
128            "avec-provider"
129        }
130
131        fn capabilities(&self) -> &'static [AiCapability] {
132            &[AiCapability::AvecScoring]
133        }
134
135        async fn embed_semantic(&self, _request: &EmbedRequest) -> Result<Vec<f32>> {
136            Ok(vec![1.0])
137        }
138
139        async fn embed_avec(&self, _request: &EmbedRequest) -> Result<Vec<f32>> {
140            Ok(vec![1.0])
141        }
142
143        async fn score_avec(&self, _request: &ScoreAvecRequest) -> Result<AvecState> {
144            Ok(AvecState {
145                stability: 0.7,
146                friction: 0.2,
147                logic: 0.9,
148                autonomy: 0.6,
149            })
150        }
151    }
152
153    #[test]
154    fn resolve_auto_picks_provider_supporting_task() {
155        let mut registry = InMemoryAiProviderRegistry::new();
156        registry.register(SemanticOnlyProvider);
157
158        let provider = registry
159            .resolve(AiTask::SemanticEmbedding, None, ProviderPolicy::Auto)
160            .expect("expected provider for semantic task");
161
162        assert_eq!(provider.provider_id(), "semantic-only");
163    }
164
165    #[test]
166    fn resolve_required_without_provider_id_fails() {
167        let mut registry = InMemoryAiProviderRegistry::new();
168        registry.register(SemanticOnlyProvider);
169
170        let err = match registry.resolve(AiTask::SemanticEmbedding, None, ProviderPolicy::Required)
171        {
172            Ok(_) => panic!("expected provider policy failure"),
173            Err(err) => err,
174        };
175
176        assert!(
177            err.to_string()
178                .contains("provider_policy=required needs an explicit provider_id")
179        );
180    }
181
182    #[test]
183    fn resolve_with_explicit_provider_requires_capability_match() {
184        let mut registry = InMemoryAiProviderRegistry::new();
185        registry.register(SemanticOnlyProvider);
186        registry.register(AvecProvider);
187
188        let err = match registry.resolve(
189            AiTask::AvecScoring,
190            Some("semantic-only"),
191            ProviderPolicy::Preferred,
192        ) {
193            Ok(_) => panic!("expected capability mismatch"),
194            Err(err) => err,
195        };
196
197        assert!(
198            err.to_string()
199                .contains("does not support task 'AvecScoring'")
200        );
201    }
202}