locus_sdk/infrastructure/
registry.rs1use 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}