1use std::collections::HashMap;
6use std::sync::Arc;
7
8use thiserror::Error;
9
10use crate::ai::cache::AiResultCache;
11use crate::ai::call_log::AiCallLog;
12use crate::ai::provider::InferenceProvider;
13use crate::ai::registry::{BackendKind, ModelBackend, ModelRegistry, RegistryError};
14
15#[derive(Clone)]
17pub struct ResolvedModel {
18 pub kind: BackendKind,
20 pub model_id: u32,
22 pub provider: Arc<dyn InferenceProvider>,
24 pub provider_model: String,
26 pub labels: Option<Vec<String>>,
28}
29
30#[derive(Debug, Error)]
32pub enum AiRuntimeError {
33 #[error(transparent)]
35 Registry(#[from] RegistryError),
36
37 #[error("model '{model}' references provider '{provider}', which is not configured")]
39 UnknownProvider {
40 model: String,
42 provider: String,
44 },
45
46 #[error("model '{0}' is local, but the local backend is not enabled in this build")]
48 LocalBackendUnavailable(String),
49}
50
51pub struct AiRuntime {
53 registry: ModelRegistry,
54 providers: HashMap<String, Arc<dyn InferenceProvider>>,
55 local_provider: Option<Arc<dyn InferenceProvider>>,
56 cache: Arc<AiResultCache>,
57 call_log: Arc<AiCallLog>,
58 model_ids: HashMap<String, u32>,
59}
60
61impl AiRuntime {
62 #[must_use]
65 pub fn new(
66 registry: ModelRegistry,
67 providers: impl IntoIterator<Item = (String, Arc<dyn InferenceProvider>)>,
68 local_provider: Option<Arc<dyn InferenceProvider>>,
69 cache: Arc<AiResultCache>,
70 call_log: Arc<AiCallLog>,
71 ) -> Self {
72 let model_ids = registry
73 .iter()
74 .enumerate()
75 .map(|(i, entry)| (entry.id.clone(), u32::try_from(i).unwrap_or(u32::MAX)))
76 .collect();
77 Self {
78 registry,
79 providers: providers.into_iter().collect(),
80 local_provider,
81 cache,
82 call_log,
83 model_ids,
84 }
85 }
86
87 #[must_use]
89 pub fn registry(&self) -> &ModelRegistry {
90 &self.registry
91 }
92
93 #[must_use]
95 pub fn cache(&self) -> &Arc<AiResultCache> {
96 &self.cache
97 }
98
99 #[must_use]
101 pub fn call_log(&self) -> &Arc<AiCallLog> {
102 &self.call_log
103 }
104
105 pub fn resolve(&self, model_name: &str) -> Result<ResolvedModel, AiRuntimeError> {
112 let entry = self.registry.resolve(model_name)?;
113 let model_id = self.model_ids.get(model_name).copied().unwrap_or(u32::MAX);
114 match &entry.backend {
115 ModelBackend::Remote { provider, model } => {
116 let client = self.providers.get(provider).ok_or_else(|| {
117 AiRuntimeError::UnknownProvider {
118 model: model_name.to_string(),
119 provider: provider.clone(),
120 }
121 })?;
122 Ok(ResolvedModel {
123 kind: BackendKind::Remote,
124 model_id,
125 provider: Arc::clone(client),
126 provider_model: model.clone(),
127 labels: None,
128 })
129 }
130 ModelBackend::Local { labels, source } => {
131 let client = self.local_provider.as_ref().ok_or_else(|| {
132 AiRuntimeError::LocalBackendUnavailable(model_name.to_string())
133 })?;
134 Ok(ResolvedModel {
135 kind: BackendKind::Local,
136 model_id,
137 provider: Arc::clone(client),
138 provider_model: source.clone(),
139 labels: labels.clone(),
140 })
141 }
142 }
143 }
144}
145
146#[cfg(test)]
147mod tests {
148 use super::*;
149 use crate::ai::provider::{
150 InferenceOutputs, InferenceRequest, InferenceResponse, ProviderError, Usage,
151 };
152 use crate::ai::registry::{ModelEntry, Task};
153 use async_trait::async_trait;
154
155 struct Stub;
156
157 #[async_trait]
158 impl InferenceProvider for Stub {
159 async fn infer_batch(
160 &self,
161 request: InferenceRequest,
162 ) -> Result<InferenceResponse, ProviderError> {
163 Ok(InferenceResponse {
164 outputs: InferenceOutputs::Text(vec![String::new(); request.inputs.len()]),
165 usage: Usage::ZERO,
166 })
167 }
168 fn name(&self) -> &'static str {
169 "stub"
170 }
171 }
172
173 fn runtime() -> AiRuntime {
174 let mut registry = ModelRegistry::new();
175 registry
176 .register(ModelEntry {
177 id: "haiku".into(),
178 tasks: vec![Task::Classify],
179 backend: ModelBackend::Remote {
180 provider: "anthropic".into(),
181 model: "claude-haiku-4-5-20251001".into(),
182 },
183 })
184 .unwrap();
185 registry
186 .register(ModelEntry {
187 id: "finbert".into(),
188 tasks: vec![Task::Classify],
189 backend: ModelBackend::Local {
190 source: "hf:onnx-community/finbert".into(),
191 labels: Some(vec!["positive".into(), "negative".into()]),
192 },
193 })
194 .unwrap();
195 let mut providers: HashMap<String, Arc<dyn InferenceProvider>> = HashMap::new();
196 providers.insert("anthropic".into(), Arc::new(Stub));
197 AiRuntime::new(
198 registry,
199 providers,
200 None,
201 Arc::new(AiResultCache::with_defaults()),
202 Arc::new(AiCallLog::with_defaults()),
203 )
204 }
205
206 #[test]
207 fn resolves_remote_model_to_its_provider() {
208 let rt = runtime();
209 let resolved = rt.resolve("haiku").unwrap();
210 assert_eq!(resolved.kind, BackendKind::Remote);
211 assert_eq!(resolved.provider.name(), "stub");
212 assert_eq!(resolved.provider_model, "claude-haiku-4-5-20251001");
213 }
214
215 #[test]
216 fn local_model_without_backend_errors() {
217 let rt = runtime();
218 assert!(matches!(
219 rt.resolve("finbert"),
220 Err(AiRuntimeError::LocalBackendUnavailable(_))
221 ));
222 }
223
224 #[test]
225 fn unknown_model_errors() {
226 let rt = runtime();
227 assert!(matches!(
228 rt.resolve("ghost"),
229 Err(AiRuntimeError::Registry(RegistryError::UnknownModel(_)))
230 ));
231 }
232}