Skip to main content

laminar_db/ai/
runtime.rs

1//! Assembled AI subsystem: registry, provider clients, result cache, and call log.
2//! Built once at startup; [`AiRuntime::resolve`] returns everything the inference
3//! operator needs to run a named model.
4
5use 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/// Everything the inference operator needs to run one model.
16#[derive(Clone)]
17pub struct ResolvedModel {
18    /// Backend kind (selects the adapter path).
19    pub kind: BackendKind,
20    /// Stable per-run integer for the result-cache key.
21    pub model_id: u32,
22    /// Transport client.
23    pub provider: Arc<dyn InferenceProvider>,
24    /// Provider-side model identifier.
25    pub provider_model: String,
26    /// Intrinsic labels for local classifiers.
27    pub labels: Option<Vec<String>>,
28}
29
30/// Errors from resolving a model to a runnable backend.
31#[derive(Debug, Error)]
32pub enum AiRuntimeError {
33    /// The model is not registered or cannot perform the requested task.
34    #[error(transparent)]
35    Registry(#[from] RegistryError),
36
37    /// A remote model references an unconfigured provider.
38    #[error("model '{model}' references provider '{provider}', which is not configured")]
39    UnknownProvider {
40        /// The model name.
41        model: String,
42        /// The missing provider name.
43        provider: String,
44    },
45
46    /// Local model referenced but the local backend is not enabled.
47    #[error("model '{0}' is local, but the local backend is not enabled in this build")]
48    LocalBackendUnavailable(String),
49}
50
51/// The assembled AI subsystem.
52pub 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    /// Assemble a runtime. `providers` is keyed by provider name; `local_provider`
63    /// serves all local models. Each model gets a stable cache id.
64    #[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    /// Model registry (backs `laminar.models`).
88    #[must_use]
89    pub fn registry(&self) -> &ModelRegistry {
90        &self.registry
91    }
92
93    /// Shared result cache.
94    #[must_use]
95    pub fn cache(&self) -> &Arc<AiResultCache> {
96        &self.cache
97    }
98
99    /// Call log (backs `laminar.ai_calls`).
100    #[must_use]
101    pub fn call_log(&self) -> &Arc<AiCallLog> {
102        &self.call_log
103    }
104
105    /// Resolve a model name to a runnable backend.
106    ///
107    /// # Errors
108    ///
109    /// Returns [`AiRuntimeError`] if the model is unknown, the provider is not
110    /// configured, or the local backend is unavailable.
111    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}